26.9(安全审计修复版)

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 7f060dd0e4
173 changed files with 32455 additions and 0 deletions
File diff suppressed because it is too large Load Diff
+667
View File
@@ -0,0 +1,667 @@
package api
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// chunkExpireTTL 分片会话保留时长(M5:预留窗口由 24h 缩短为 2h;
// 会话本身保留 24h 支持断点续传,见 janitor 的清理周期)。
const chunkExpireTTL = 2 * time.Hour
// maxChunkSizeBytes 单分片大小上限 32MBM3:限制 io.ReadAll 内存占用)。
const maxChunkSizeBytes = 32 * 1024 * 1024
// ============ POST /chunk/upload/init 初始化分片会话 ============
// requireChunkEnabled L4enableChunk 开关后端强制(此前仅前端隐藏入口,
// 开关关闭后 /chunk/* 接口仍可直接调用)。
func (d *Deps) requireChunkEnabled(c *gin.Context) bool {
if d.Cfg.EnableChunk() {
return true
}
auditRecordFailed(c, d.AuditSvc, "分片上传未启用")
response.Fail(c, http.StatusForbidden, "分片上传未启用")
return false
}
// chunkInitRequest init 请求体(JSON 或表单)。
type chunkInitRequest struct {
FileName string `json:"file_name" form:"file_name"`
ChunkSize int64 `json:"chunk_size" form:"chunk_size"`
FileSize int64 `json:"file_size" form:"file_size"`
FileHash string `json:"file_hash" form:"file_hash"`
}
// chunkInit 创建分片上传会话(对齐参考 init_chunk_upload):
// 支持断点续传(相同 hash/大小/文件名的未完成会话直接续传)。
func (d *Deps) chunkInit(c *gin.Context) {
if !d.requireChunkEnabled(c) {
return
}
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
var req chunkInitRequest
if err := bindJSONOrForm(c, &req); err != nil {
respondError(c, err)
return
}
safeName := storage.SanitizeFileName(req.FileName)
if safeName == "" {
auditRecordFailed(c, d.AuditSvc, "文件名非法")
response.Fail(c, http.StatusBadRequest, "文件名非法")
return
}
// 文件类型白名单(无内容可校验,仅名称)
if err := validateFileMagic(d.Cfg, safeName, "", nil); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件类型被拒绝")
respondError(c, err)
return
}
chunkSize := req.ChunkSize
if chunkSize <= 0 {
chunkSize = 5 * 1024 * 1024 // 默认 5MB(对齐参考 InitChunkUploadModel
}
// M3:单片全部读入内存后再落存储,必须限制单片大小(客户端声明的
// chunk_size 上界受策略约束,但策略允许至 10GiB → 显式封顶 32MB)。
if chunkSize > maxChunkSizeBytes {
auditRecordFailed(c, d.AuditSvc, "chunk_size 超过上限")
response.Fail(c, http.StatusBadRequest, fmt.Sprintf("chunk_size 过大,最大为 %d MB", maxChunkSizeBytes>>20))
return
}
if req.FileSize <= 0 {
auditRecordFailed(c, d.AuditSvc, "file_size 非法")
response.Fail(c, http.StatusBadRequest, "file_size 必须大于 0")
return
}
// 服务端按分片数上限校验总大小(防分片声明绕过)
totalChunks := (req.FileSize + chunkSize - 1) / chunkSize
maxPossible := totalChunks * chunkSize
// v2 需求 ④⑩:动态策略校验(max_file_size0=回落 uploadSize
if err := d.CurrentUploadPolicy().CheckSize(maxPossible); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小超过限制")
respondError(c, err)
return
}
ctx := c.Request.Context()
// 断点续传:查找相同 hash+大小+文件名的未完成会话(chunk_index=-1 为会话头)
var existing model.UploadChunk
err := d.DB.WithContext(ctx).
Where("chunk_hash = ? AND chunk_index = -1 AND file_size = ? AND file_name = ?",
req.FileHash, req.FileSize, safeName).
First(&existing).Error
if err == nil {
if existing.SavePath == "" {
// 脏会话:清理后按新建处理
_ = d.DB.WithContext(ctx).
Where("upload_id = ?", existing.UploadID).
Delete(&model.UploadChunk{}).Error
releaseStorage(ctx, d.DB, "chunk:"+existing.UploadID)
} else {
if err := reserveStorage(ctx, d.DB, d.Cfg, "chunk:"+existing.UploadID, existing.FileSize, chunkExpireTTL); err != nil {
respondError(c, err)
return
}
uploaded := d.uploadedChunkIndexes(ctx, existing.UploadID)
auditUploadEntry(c, existing.UploadID, safeName, req.FileSize, 0)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{
"existed": false,
"upload_id": existing.UploadID,
"chunk_size": existing.ChunkSize,
"total_chunks": existing.TotalChunks,
"uploaded_chunks": uploaded,
})
return
}
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
// 新建会话
uploadID := uuidHex()
resToken := "chunk:" + uploadID
if err := reserveStorage(ctx, d.DB, d.Cfg, resToken, req.FileSize, chunkExpireTTL); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "容量预留失败")
respondError(c, err)
return
}
// M5:init 即计入上传限流(此前仅 complete 成功时计数,
// 恶意客户端可无限创建会话占用容量预留)
d.Limiter.Add(c, middleware.LimitUpload)
_, _, _, _, savePath := buildSavePath(d.Cfg, safeName, uploadID)
session := model.UploadChunk{
UploadID: uploadID,
ChunkIndex: -1,
TotalChunks: int(totalChunks),
FileSize: req.FileSize,
ChunkSize: int(chunkSize),
ChunkHash: req.FileHash,
FileName: safeName,
SavePath: savePath,
Engine: d.Store.CurrentName(), // v3:会话归属引擎(分片/合并全程走同一引擎)
}
if err := d.DB.WithContext(ctx).Create(&session).Error; err != nil {
releaseStorage(ctx, d.DB, resToken)
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "会话创建失败")
respondError(c, errInternal("创建上传会话失败: "+err.Error()))
return
}
auditUploadEntry(c, uploadID, safeName, req.FileSize, 0)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{
"existed": false,
"upload_id": uploadID,
"chunk_size": chunkSize,
"total_chunks": totalChunks,
"uploaded_chunks": []int{},
})
}
// uploadedChunkIndexes 查询会话中已完成分片的索引列表。
func (d *Deps) uploadedChunkIndexes(ctx context.Context, uploadID string) []int {
var rows []model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND completed = ?", uploadID, true).
Order("chunk_index ASC").Find(&rows).Error; err != nil {
return []int{}
}
out := make([]int, 0, len(rows))
for _, r := range rows {
out = append(out, r.ChunkIndex)
}
return out
}
// ============ POST /chunk/upload/{uploadID}/{index}(及扁平兼容)============
// chunkUploadFlat 扁平模式:POST /chunk/uploadupload_id/chunk_index 走表单或 query。
// 多文件字段(chunk/chunks)时按 base_chunk_index 顺序批量接收。
func (d *Deps) chunkUploadFlat(c *gin.Context) {
c.Params = append(c.Params, gin.Param{Key: "uploadID", Value: resolveUploadID(c)})
c.Params = append(c.Params, gin.Param{Key: "chunkIndex", Value: resolveChunkIndex(c)})
d.chunkUpload(c)
}
// resolveUploadID 解析 upload_id:路径参数 → multipart 表单 → query。
func resolveUploadID(c *gin.Context) string {
if v := c.Param("uploadID"); v != "" {
return v
}
if v := c.PostForm("upload_id"); v != "" {
return v
}
return c.Query("upload_id")
}
// resolveChunkIndex 解析 chunk_index:路径参数 → multipart 表单 → query。
func resolveChunkIndex(c *gin.Context) string {
if v := c.Param("chunkIndex"); v != "" {
return v
}
if v := c.PostForm("chunk_index"); v != "" {
return v
}
return c.Query("chunk_index")
}
// chunkUpload 上传单个(或批量)分片(对齐参考 upload_chunk)。
// multipart 文件字段:chunk(主)或 file(回退);批量用 chunk[]/chunks 数组 + chunk_index 为起始索引。
func (d *Deps) chunkUpload(c *gin.Context) {
if !d.requireChunkEnabled(c) {
return
}
if !d.requireShareLogin(c) {
return
}
uploadID := resolveUploadID(c)
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, "上传会话不存在")
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
if err := reserveStorage(ctx, d.DB, d.Cfg, "chunk:"+uploadID, session.FileSize, chunkExpireTTL); err != nil {
respondError(c, err)
return
}
// 收集分片文件:chunk(单)→ file(回退)→ chunk[]/chunks(批量)
form, err := c.MultipartForm()
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "multipart 解析失败")
response.Fail(c, http.StatusBadRequest, "multipart 表单解析失败")
return
}
files := form.File["chunk"]
single := len(files) == 0
if single {
files = form.File["file"]
}
if len(files) == 0 {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "缺少 chunk 分片字段")
response.Fail(c, http.StatusBadRequest, "缺少分片文件字段 chunk")
return
}
baseIndex, err := strconv.Atoi(resolveChunkIndex(c))
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "无效的分片索引")
response.Fail(c, http.StatusBadRequest, "无效的分片索引")
return
}
results := make([]gin.H, 0, len(files))
for i, fh := range files {
// 单分片模式严格使用请求索引;批量模式从 base 递增
idx := baseIndex
if !single && len(files) > 1 {
idx = baseIndex + i
}
res, status, msg := d.saveOneChunk(c, ctx, &session, idx, fh)
if status != 0 {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, msg)
response.Fail(c, status, msg)
return
}
results = append(results, res)
}
// 审计:传输字节数为本次请求分片总和
var transferred int64
for _, fh := range files {
transferred += fh.Size
}
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, transferred)
auditRecordSuccess(c, d.AuditSvc)
if len(results) == 1 {
response.OK(c, results[0])
return
}
response.OK(c, gin.H{"chunks": results})
}
// saveOneChunk 保存一个分片:查重→读数据→校验→存储→记录。
// 返回 (响应体, HTTP错误状态码, 错误信息);成功时状态码为 0。
func (d *Deps) saveOneChunk(c *gin.Context, ctx context.Context, session *model.UploadChunk, idx int, fh *multipart.FileHeader) (gin.H, int, string) {
if idx < 0 || idx >= session.TotalChunks {
return nil, http.StatusBadRequest, "无效的分片索引"
}
// 已上传分片:断点续传直接跳过
var existing model.UploadChunk
err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = ? AND completed = ?", session.UploadID, idx, true).
First(&existing).Error
if err == nil {
return gin.H{"chunk_hash": existing.ChunkHash, "skipped": true, "chunk_index": idx}, 0, ""
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, http.StatusInternalServerError, "查询分片记录失败"
}
f, err := fh.Open()
if err != nil {
return nil, http.StatusBadRequest, "分片数据读取失败"
}
defer func() { _ = f.Close() }()
data, err := io.ReadAll(io.LimitReader(f, int64(session.ChunkSize)+1))
if err != nil {
return nil, http.StatusBadRequest, "分片数据读取失败"
}
// 校验分片大小不超过声明值
if int64(len(data)) > int64(session.ChunkSize) {
return nil, http.StatusBadRequest,
"分片大小超过声明值: 最大 " + strconv.Itoa(session.ChunkSize) + ", 实际 " + strconv.Itoa(len(data))
}
// 累计大小校验(已传分片数×chunk_size + 当前分片;动态策略上限)
var uploadedCount int64
_ = d.DB.WithContext(ctx).Model(&model.UploadChunk{}).
Where("upload_id = ? AND completed = ?", session.UploadID, true).
Count(&uploadedCount).Error
if err := d.CurrentUploadPolicy().CheckSize(uploadedCount*int64(session.ChunkSize) + int64(len(data))); err != nil {
return nil, http.StatusForbidden, err.Error()
}
// 首分片做 magic bytes 防伪造
if idx == 0 {
head := data
if len(head) > 64 {
head = head[:64]
}
if err := validateFileMagic(d.Cfg, session.FileName, "", head); err != nil {
return nil, http.StatusForbidden, "文件内容校验失败:" + err.Error()
}
}
sum := sha256.Sum256(data)
chunkHash := hex.EncodeToString(sum[:])
if _, err := d.Store.SaveChunk(ctx, session.UploadID, idx, bytes.NewReader(data), session.SavePath); err != nil {
return nil, http.StatusInternalServerError, "分片保存失败: " + err.Error()
}
// 保存成功后再记录(对齐参考:先存储后落库)。
// 注意:不能用结构体 Where 条件(GORM 会忽略零值字段,chunk_index=0 会被
// 丢弃从而误匹配 -1 会话行),必须用字符串条件 + 完整目标结构体。
rec := model.UploadChunk{
UploadID: session.UploadID,
ChunkIndex: idx,
ChunkHash: chunkHash,
Completed: true,
FileSize: session.FileSize,
TotalChunks: session.TotalChunks,
ChunkSize: session.ChunkSize,
FileName: session.FileName,
SavePath: session.SavePath,
Engine: session.Engine, // v3:继承会话引擎
}
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = ?", session.UploadID, idx).
FirstOrCreate(&rec).Error; err != nil {
return nil, http.StatusInternalServerError, "分片记录写入失败"
}
return gin.H{"chunk_hash": chunkHash, "chunk_index": idx}, 0, ""
}
// ============ GET /chunk/upload/status/{uploadID} ============
// chunkStatus 查询上传进度(对齐参考 get_upload_status)。
func (d *Deps) chunkStatus(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
if uploadID == "" {
uploadID = c.Query("upload_id")
}
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
uploaded := d.uploadedChunkIndexes(ctx, uploadID)
var progress float64
if session.TotalChunks > 0 {
progress = float64(len(uploaded)) / float64(session.TotalChunks) * 100
}
response.OK(c, gin.H{
"upload_id": uploadID,
"file_name": session.FileName,
"file_size": session.FileSize,
"chunk_size": session.ChunkSize,
"total_chunks": session.TotalChunks,
"uploaded_chunks": uploaded,
"progress": progress,
})
}
// ============ POST /chunk/upload/complete/{uploadID} ============
// chunkCompleteRequest complete 请求体。
type chunkCompleteRequest struct {
ExpireValue int `json:"expire_value" form:"expire_value"`
ExpireStyle string `json:"expire_style" form:"expire_style"`
Code string `json:"code" form:"code"` // v3.1:自定义提取码(4-8 位字母数字,空=随机)
}
// chunkComplete 合并分片并创建分享(对齐参考 complete_upload)。
func (d *Deps) chunkComplete(c *gin.Context) {
if !d.requireChunkEnabled(c) {
return
}
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
uploadID := c.Param("uploadID")
if uploadID == "" {
uploadID = resolveUploadID(c)
}
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, "上传会话不存在")
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
var req chunkCompleteRequest
if err := bindJSONOrForm(c, &req); err != nil {
respondError(c, err)
return
}
exp, err := resolveExpire(d.Cfg, req.ExpireValue, req.ExpireStyle)
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "过期策略非法")
respondError(c, err)
return
}
// v3.1:自定义提取码(合并前校验,失败快速返回)
if err := validatePickupCode(req.Code); err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
if err := reserveStorage(ctx, d.DB, d.Cfg, "chunk:"+uploadID, session.FileSize, chunkExpireTTL); err != nil {
respondError(c, err)
return
}
// 分片完整性校验(chunk_index >= 0-1 为会话头,completed 恒为 false
var completed []model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND completed = ? AND chunk_index >= 0", uploadID, true).
Find(&completed).Error; err != nil {
respondError(c, errInternal("查询分片记录失败: "+err.Error()))
return
}
if len(completed) != session.TotalChunks {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "分片不完整")
response.Fail(c, http.StatusBadRequest, "分片不完整")
return
}
// 累计大小上限校验(超限清理会话,对齐参考;动态策略上限)
if err := d.CurrentUploadPolicy().CheckSize(int64(len(completed)) * int64(session.ChunkSize)); err != nil {
if cs, ce := d.storeFor(session.Engine); ce == nil {
_ = cs.CleanChunks(ctx, uploadID, session.SavePath)
}
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).Delete(&model.UploadChunk{}).Error
releaseStorage(ctx, d.DB, "chunk:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "实际上传大小超过限制")
respondError(c, err)
return
}
// 合并(引擎负责按索引有序合并+SHA256 校验)
verifyHash := func(index int) (string, error) {
var rec model.UploadChunk
err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = ?", uploadID, index).
First(&rec).Error
if err != nil {
return "", err
}
return rec.ChunkHash, nil
}
// v3:合并走会话归属引擎(会话创建时的引擎,即使中途热切换也不受影响)
mergeStore, sErr := d.storeFor(session.Engine)
if sErr != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, session.FileSize)
auditRecordFailed(c, d.AuditSvc, "存储引擎不可用: "+sErr.Error())
respondError(c, mapStorageError(sErr))
return
}
size, fileHash, err := mergeStore.MergeChunks(ctx, uploadID, session.TotalChunks, verifyHash, session.SavePath)
if err != nil {
_ = mergeStore.CleanChunks(ctx, uploadID, session.SavePath)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, session.FileSize)
auditRecordFailed(c, d.AuditSvc, "文件合并失败")
respondError(c, mapStorageError(err))
return
}
// 创建分享记录(v3.1:支持自定义提取码)
code, err := pickCustomCode(ctx, d.DB, d.Cfg, req.Code)
if err == nil {
fc := model.FileCodes{
Code: code,
FileHash: &fileHash,
IsChunked: true,
UploadID: &uploadID,
Size: session.FileSize,
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: session.Engine, // v3:归属引擎戳
}
// 拆分路径与文件名(对齐参考:path=dirname(save_path), uuid=basename
dir, name := splitDirBase(session.SavePath)
ext := baseExt(name)
fc.FilePath = &dir
fc.UUIDFileName = &name
fc.Prefix = trimExt(name)
fc.Suffix = ext
err = d.DB.WithContext(ctx).Create(&fc).Error
err = mapCodeConflict(err) // v3.1
}
if err == nil {
// 成功:清理分片与记录(走归属引擎)
_ = mergeStore.CleanChunks(ctx, uploadID, session.SavePath)
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).Delete(&model.UploadChunk{}).Error
}
releaseStorage(ctx, d.DB, "chunk:"+uploadID)
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, size)
auditRecordFailed(c, d.AuditSvc, "创建分享失败")
respondError(c, errInternal("创建分享失败: "+err.Error()))
return
}
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, session.FileName, session.FileSize, size)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": session.FileName})
}
// splitDirBase 拆分相对路径为目录与文件名。
func splitDirBase(p string) (dir, base string) {
for i := len(p) - 1; i >= 0; i-- {
if p[i] == '/' {
return p[:i], p[i+1:]
}
}
return "", p
}
// baseExt 提取扩展名(含点)。
func baseExt(name string) string {
for i := len(name) - 1; i >= 0; i-- {
if name[i] == '.' {
return name[i:]
}
if name[i] == '/' {
break
}
}
return ""
}
// trimExt 去除扩展名。
func trimExt(name string) string {
ext := baseExt(name)
return name[:len(name)-len(ext)]
}
// ============ DELETE /chunk/upload/{uploadID} 取消上传 ============
// chunkCancel 取消上传并清理临时文件(对齐参考 cancel_upload)。
func (d *Deps) chunkCancel(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
if uploadID == "" {
uploadID = c.Query("upload_id")
}
ctx := c.Request.Context()
var session model.UploadChunk
if err := d.DB.WithContext(ctx).
Where("upload_id = ? AND chunk_index = -1", uploadID).
First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
if session.SavePath != "" {
if cs, ce := d.storeFor(session.Engine); ce == nil {
_ = cs.CleanChunks(ctx, uploadID, session.SavePath)
}
}
if err := d.DB.WithContext(ctx).
Where("upload_id = ?", uploadID).
Delete(&model.UploadChunk{}).Error; err != nil {
respondError(c, errInternal("取消上传失败: "+err.Error()))
return
}
releaseStorage(ctx, d.DB, "chunk:"+uploadID)
response.OK(c, gin.H{"message": "上传已取消"})
}
+798
View File
@@ -0,0 +1,798 @@
// Package api 提供文件快传的 HTTP API 层:
// 分享(share)、分片上传(chunk)、预签名直传(presign)、管理端(admin)、
// 初始化向导(setup)与前端静态资源(web)。
//
// 接口语义对齐参考实现 apps/base/views.py 与 apps/admin/views.py
// 响应统一 {"code":200,"msg":"","data":...}internal/response)。
// 所有上传/下载端点经 middleware.Audit 落审计日志(需求 ③)。
package api
import (
"bytes"
"context"
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"math/big"
"mime/multipart"
"net/http"
"net/url"
"path"
"regexp"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/audit"
"filecodebox/internal/config"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// apiError 统一的业务错误:handler 返回该错误并由 respondError 映射 HTTP 状态。
type apiError struct {
Status int // HTTP 状态码(同时写入响应体 code)
Msg string // 错误描述(中文)
}
func (e *apiError) Error() string { return e.Msg }
func errBadRequest(msg string) error { return &apiError{Status: http.StatusBadRequest, Msg: msg} }
func errForbidden(msg string) error { return &apiError{Status: http.StatusForbidden, Msg: msg} }
func errNotFound(msg string) error { return &apiError{Status: http.StatusNotFound, Msg: msg} }
func errConflict(msg string) error { return &apiError{Status: http.StatusConflict, Msg: msg} }
func errLocked(msg string) error { return &apiError{Status: http.StatusLocked, Msg: msg} }
func errInternal(msg string) error {
return &apiError{Status: http.StatusInternalServerError, Msg: msg}
}
func errInsufficient(msg string) error {
return &apiError{Status: http.StatusInsufficientStorage, Msg: msg}
}
func errNotImpl(msg string) error { return &apiError{Status: http.StatusNotImplemented, Msg: msg} }
// mapStorageError 把存储层哨兵错误映射为 apiError(对齐 README 约定):
// ErrNotFound→404、ErrInvalidPath→400、ErrUnavailable→503、ErrNotSupported→501、
// ErrRangeNotSatisfiable→416、ErrHashMismatch→400(支持 %w 包装判定)。
func mapStorageError(err error) error {
if err == nil {
return nil
}
switch {
case errors.Is(err, storage.ErrNotFound):
return errNotFound("文件不存在")
case errors.Is(err, storage.ErrInvalidPath):
return errBadRequest("非法文件路径")
case errors.Is(err, storage.ErrUnavailable):
return &apiError{Status: http.StatusServiceUnavailable, Msg: "存储服务不可用,请稍后再试"}
case errors.Is(err, storage.ErrNotSupported):
return errNotImpl("当前存储引擎不支持该操作")
case errors.Is(err, storage.ErrRangeNotSatisfiable):
return &apiError{Status: http.StatusRequestedRangeNotSatisfiable, Msg: "请求范围超出文件大小"}
case errors.Is(err, storage.ErrHashMismatch):
return errBadRequest("分片哈希校验失败,请重新上传")
}
return errInternal("存储操作失败: " + err.Error())
}
// respondError 统一错误出口:apiError 按其状态码响应,其余 500。
func respondError(c *gin.Context, err error) {
var ae *apiError
if !errors.As(err, &ae) {
ae = &apiError{Status: http.StatusInternalServerError, Msg: err.Error()}
}
// 非 2xx 且未标记审计跳过时,兜底给出 result/errorMsg(中间件还会按状态兜底)
middleware.AuditSet(c, func(e *audit.Entry) {
if e.Result == "" {
if ae.Status == 401 || ae.Status == 403 || ae.Status == 423 || ae.Status == 429 || ae.Status == 428 {
e.Result = model.AuditResultDenied
} else {
e.Result = model.AuditResultFailed
}
}
if e.ErrorMsg == "" {
e.ErrorMsg = ae.Msg
}
})
response.Fail(c, ae.Status, ae.Msg)
}
// ============ 取件码 / 下载令牌 ============
// codeChars 随机字符码字符集(对齐参考 string.ascii_uppercase + digits)。
const codeChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
// generateCode 生成取件码:secret→5 位大写字母+数字;number→5 位数字(10000~99999)。
// 与参考 get_random_code 一致:冲突时重试(由调用方控制上限)。
func generateCode(style string) string {
if style == "number" {
n, _ := rand.Int(rand.Reader, big.NewInt(90000))
return strconv.FormatInt(10000+n.Int64(), 10)
}
b := make([]byte, 5)
for i := range b {
n, _ := rand.Int(rand.Reader, big.NewInt(int64(len(codeChars))))
b[i] = codeChars[n.Int64()]
}
return string(b)
}
// ============ 自定义提取码(v3.1,防撞库)============
const (
// pickupCodeMinLen 最小长度:L3 由 4 提升至 5(4 位码空间仅 168 万,
// 多 IP 分布式撞库可在限流窗口内覆盖;5 位起 6000 万空间显著提高成本)
pickupCodeMinLen = 5
pickupCodeMaxLen = 8
)
// validatePickupCode 校验自定义提取码:
// - 长度 5-8 字符(下限保证码空间 36^5≈6000 万起,配合取件错误限流防撞库);
// - 仅字母/数字(排除空格与符号,规避中文场景误输与 URL 转义问题);
// - 空串合法(空 = 使用系统随机码)。
// - 兼容:仅约束新建码,历史 4 位码仍可正常取件。
func validatePickupCode(code string) error {
code = strings.TrimSpace(code)
if code == "" {
return nil
}
if l := len(code); l < pickupCodeMinLen || l > pickupCodeMaxLen {
return errBadRequest(fmt.Sprintf("提取码长度须为 %d-%d 位", pickupCodeMinLen, pickupCodeMaxLen))
}
for _, r := range code {
if !(r >= '0' && r <= '9') && !(r >= 'a' && r <= 'z') && !(r >= 'A' && r <= 'Z') {
return errBadRequest("提取码仅支持字母和数字")
}
}
return nil
}
// pickCustomCode 解析提取码:自定义非空→查重返回;空→随机码。
// 占用时 400(不静默改码——用户可能已把链接发出去,静默改码会导致取不到件)。
func pickCustomCode(ctx context.Context, db *gorm.DB, cfg *config.Config, custom string) (string, error) {
custom = strings.ToUpper(strings.TrimSpace(custom))
if custom == "" {
return randomCode(ctx, db, cfg)
}
var count int64
if err := db.WithContext(ctx).Model(&model.FileCodes{}).Where("code = ?", custom).Count(&count).Error; err != nil {
return "", errInternal("取件码查重失败: " + err.Error())
}
if count > 0 {
return "", errBadRequest("该提取码已被占用,请换一个")
}
return custom, nil
}
// mapCodeConflict v3.1:分享记录创建失败时,若是自定义码唯一索引冲突(并发兜底,
// pickCustomCode 的预查重未覆盖竞态),转为友好 400;其余错误原样返回。
func mapCodeConflict(err error) error {
if err == nil {
return nil
}
msg := err.Error()
if strings.Contains(msg, "已被占用") ||
strings.Contains(msg, "UNIQUE constraint") ||
strings.Contains(msg, "duplicate key") ||
strings.Contains(msg, "Duplicate entry") {
return errBadRequest("该提取码已被占用,请换一个")
}
return err
}
// ============ 站点对外域名(v3.1============
// SiteDomain 站点对外域名规范化(v3.1):空串合法(分享链接用当前访问地址)。
// 接受 http(s)://host[:port] 或裸 host[:port](自动补 http://,内网场景)。
// 拒绝路径/查询/片段/用户信息/非 http(s) 协议/非法主机字符(防 javascript: 注入分享链接)。
var siteDomainHostRe = regexp.MustCompile(`^[A-Za-z0-9]([A-Za-z0-9-]*[A-Za-z0-9])?(\.[A-Za-z0-9]([A-Za-z0-9-]*[A-Za-z0-9])?)*$`)
func normalizeSiteDomain(raw string) (string, error) {
s := strings.TrimSpace(raw)
if s == "" {
return "", nil
}
s = strings.TrimRight(s, "/")
if !strings.Contains(s, "://") {
s = "http://" + s
}
u, err := url.Parse(s)
if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") {
return "", errBadRequest("站点域名格式:http(s)://主机[:端口],如 https://share.example.com")
}
if (u.Path != "" && u.Path != "/") || u.RawQuery != "" || u.Fragment != "" || u.User != nil {
return "", errBadRequest("站点域名只填主机与端口,不带路径,如 https://share.example.com")
}
if !siteDomainHostRe.MatchString(u.Hostname()) {
return "", errBadRequest("站点域名主机名仅支持字母、数字、点与连字符")
}
if p := u.Port(); p != "" {
if n, perr := strconv.Atoi(p); perr != nil || n < 1 || n > 65535 {
return "", errBadRequest("站点域名端口须为 1-65535")
}
}
return s, nil
}
// randomCode 生成唯一取件码(查库去重,最多尝试 20 次防止极端碰撞)。
// style 为空时取配置 code_generate_typesecret/string→secret,其余→number)。
func randomCode(ctx context.Context, db *gorm.DB, cfg *config.Config) (string, error) {
style := strings.TrimSpace(cfg.GetString("code_generate_type"))
if style == "string" {
style = "secret"
}
if style != "secret" && style != "number" {
style = "number"
}
for i := 0; i < 20; i++ {
code := generateCode(style)
var count int64
if err := db.WithContext(ctx).Model(&model.FileCodes{}).
Where("code = ?", code).Count(&count).Error; err != nil {
return "", errInternal("取件码生成失败: " + err.Error())
}
if count == 0 {
return code, nil
}
}
return "", errInternal("取件码生成失败,请重试")
}
// GetSelectToken 生成下载令牌(L2HMAC-SHA256 替换拼接哈希,消除拼接歧义;
// 密钥前置为 HMAC key,窗口语义不变):
// HMAC-SHA256(key=secret, msg=code|time_factor)time_factor = unix秒/1000 - offset。
// offset=0 当前窗口、offset=1 上一窗口——下载端点同时接受两个窗口,
// 避免 ~16.7 分钟窗口边界竞态导致偶发 403。
func GetSelectToken(code, secret string, offset int) string {
timeFactor := time.Now().Unix()/1000 - int64(offset)
mac := hmac.New(sha256.New, []byte(secret))
fmt.Fprintf(mac, "%s|%d", code, timeFactor)
return hex.EncodeToString(mac.Sum(nil))
}
// VerifySelectToken 常量时间校验下载令牌(当前与上一窗口任一匹配即通过)。
func VerifySelectToken(code, secret, key string) bool {
return hmac.Equal([]byte(key), []byte(GetSelectToken(code, secret, 0))) ||
hmac.Equal([]byte(key), []byte(GetSelectToken(code, secret, 1)))
}
// ============ 过期策略 ============
// expireResult 过期策略解析结果(对齐参考 get_expire_info)。
type expireResult struct {
ExpiredAt *time.Time // nil 表示永久
ExpiredCount int // <0 按时间过期;>0 按次数
UsedCount int
}
// resolveExpire 校验 expire_style 白名单并计算过期信息。
// 对齐参考:max_save_seconds>0 时为最长保存上限(超限 403),否则默认 7 天上限;
// v2 需求 ④:style=count 时 expire_value 不得超出 max_save_count0=不限制,超限 403)。
func resolveExpire(cfg *config.Config, expireValue int, expireStyle string) (*expireResult, error) {
allowed := cfg.ExpireStyle()
okStyle := false
for _, s := range allowed {
if s == expireStyle {
okStyle = true
break
}
}
if !okStyle {
return nil, errBadRequest("过期时间类型错误")
}
if expireValue <= 0 {
return nil, errBadRequest("过期时间值必须大于 0")
}
now := time.Now()
res := &expireResult{ExpiredCount: -1, UsedCount: 0}
var expiredAt time.Time
switch expireStyle {
case "day":
expiredAt = now.AddDate(0, 0, expireValue)
case "hour":
expiredAt = now.Add(time.Duration(expireValue) * time.Hour)
case "minute":
expiredAt = now.Add(time.Duration(expireValue) * time.Minute)
case "count":
// 保存次数策略(需求 ④):max_save_count>0 时为可取次数上限,超限 403
if maxCount := cfg.MaxSaveCount(); maxCount > 0 && expireValue > maxCount {
return nil, errForbidden(fmt.Sprintf("限制次数最多为 %d 次", maxCount))
}
// 按次数过期:固定保留 1 天时间兜底(对齐参考)
expiredAt = now.AddDate(0, 0, 1)
res.ExpiredCount = expireValue
case "forever":
res.ExpiredAt = nil
res.ExpiredCount = -1
return res, nil
default:
expiredAt = now.AddDate(0, 0, 1)
}
// 最长保存时间限制
maxSeconds := cfg.MaxSaveSeconds()
maxDelta := 7 * 24 * time.Hour
if maxSeconds > 0 {
maxDelta = time.Duration(maxSeconds) * time.Second
}
if expiredAt.Sub(now) > maxDelta {
return nil, errForbidden(fmt.Sprintf("限制最长时间为 %s,可换用其他方式", formatDurationCN(maxDelta)))
}
res.ExpiredAt = &expiredAt
return res, nil
}
// formatDurationCN 把时长格式化为中文描述(对齐参考 max_save_times_desc)。
func formatDurationCN(d time.Duration) string {
sec := int64(d.Seconds())
days := sec / 86400
hours := sec % 86400 / 3600
minutes := sec % 3600 / 60
seconds := sec % 60
var parts []string
if days > 0 {
parts = append(parts, fmt.Sprintf("%d天", days))
}
if hours > 0 {
parts = append(parts, fmt.Sprintf("%d小时", hours))
}
if minutes > 0 {
parts = append(parts, fmt.Sprintf("%d分钟", minutes))
}
if seconds > 0 {
parts = append(parts, fmt.Sprintf("%d秒", seconds))
}
if len(parts) == 0 {
return "0秒"
}
return strings.Join(parts, "")
}
// ============ 存储路径 / 容量预留 ============
// storeFor v3:按归属引擎取存储实例(空戳/未知名回落当前引擎,兼容历史数据)。
func (d *Deps) storeFor(engine string) (storage.Storage, error) {
if engine == "" || !storage.ValidEngine(engine) {
return d.Store, nil
}
if engine == d.Store.CurrentName() {
return d.Store, nil
}
return d.Store.EngineOf(engine)
}
// buildSavePath 生成上传文件的存储相对路径(对齐参考 get_file_path_name):
// [storage_path/]share/data/YYYY/MM/DD/<uuid>/<清理后文件名>。
func buildSavePath(cfg *config.Config, rawName string, fileUUID string) (dirPath, prefix, suffix, cleanName, savePath string) {
today := time.Now().Format("2006/01/02")
cleanName = storage.SanitizeFileName(rawName)
ext := path.Ext(cleanName)
prefix = strings.TrimSuffix(cleanName, ext)
suffix = ext
base := "share/data/" + today + "/" + fileUUID
if sp := strings.Trim(cfg.GetString("storage_path"), "/"); sp != "" {
base = sp + "/" + base
}
dirPath = base
savePath = base + "/" + cleanName
return
}
// reserveStorage 原子预留上传容量(对齐参考 quota.reserve_storage):
// storageLimit<=0 或 size=0 时不限制直接返回;
// 通过单条 INSERT...SELECT 条件写入保证 (已用+已预留+本次) <= limit,超限返回 507。
func reserveStorage(ctx context.Context, db *gorm.DB, cfg *config.Config, token string, size int64, ttl time.Duration) error {
limit := cfg.GetInt64("storageLimit")
if limit <= 0 || size <= 0 {
return nil
}
now := time.Now()
expiresAt := now.Add(ttl)
// 清理同 token 的过期预留
if err := db.WithContext(ctx).
Where("token = ? AND expires_at <= ?", token, now).
Delete(&model.StorageReservation{}).Error; err != nil {
return errInternal("容量预留失败: " + err.Error())
}
// 同 token 已有生效预留:大小一致则幂等返回,不一致报冲突(对齐参考 409)
var existing model.StorageReservation
err := db.WithContext(ctx).
Where("token = ? AND expires_at > ?", token, now).First(&existing).Error
if err == nil {
if existing.Size == size {
return nil
}
return errConflict("上传容量预留信息不一致")
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return errInternal("容量预留失败: " + err.Error())
}
// 原子条件插入:已用 + 生效预留 + 本次 <= 上限。
// Info5Postgres READ COMMITTED 下并发 INSERT..SELECT 可能同时读到相同快照
// 而轻微超额记账,故包事务并用事务级 advisory lock 串行化配额判定
// (SQLite 写本身串行,无需加锁)。
lockFn := func(tx *gorm.DB) error {
if tx.Dialector.Name() == config.DBDriverPostgres {
return tx.Exec(`SELECT pg_advisory_xact_lock(?)`, quotaLockKey).Error
}
return nil
}
var insertErr error
txErr := db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
if err := lockFn(tx); err != nil {
return err
}
res := tx.Exec(`
INSERT INTO storage_reservations (token, size, expires_at)
SELECT ?, ?, ?
WHERE (
COALESCE((SELECT COALESCE(SUM(size),0) FROM file_codes), 0)
+ COALESCE((SELECT COALESCE(SUM(size),0) FROM storage_reservations WHERE expires_at > ?), 0)
+ ?
) <= ?`,
token, size, expiresAt, now, size, limit)
insertErr = res.Error
if res.Error != nil {
return res.Error // 触发回滚(同 token 冲突分支在外层处理)
}
if res.RowsAffected == 0 {
return errInsufficient("存储空间已达到管理员设置的容量上限")
}
return nil
})
if txErr != nil {
// 并发冲突回退:检查是否已有同 token 同大小的生效预留(对齐参考并发分支)
var cnt int64
_ = db.WithContext(ctx).Model(&model.StorageReservation{}).
Where("token = ? AND size = ? AND expires_at > ?", token, size, now).
Count(&cnt).Error
if cnt > 0 {
return nil
}
var ins *apiError
if errors.As(txErr, &ins) && ins.Status == http.StatusInsufficientStorage {
return txErr // 507:真实容量不足
}
if insertErr != nil && errors.Is(insertErr, txErr) {
return errInternal("容量预留失败: " + insertErr.Error())
}
return errInternal("容量预留失败: " + txErr.Error())
}
return nil
}
// quotaLockKey Postgres advisory lock 键(配额判定的事务级串行化)。
const quotaLockKey int64 = 0x46434251 // "FCBQ"
// releaseStorage 释放容量预留(幂等)。
func releaseStorage(ctx context.Context, db *gorm.DB, token string) {
_ = db.WithContext(ctx).Where("token = ?", token).Delete(&model.StorageReservation{}).Error
}
// ============ 文件类型校验(对齐 apps/base/file_validation.py============
// fileKind 已知文件类型:扩展名 / MIME / magic bytes。
type fileKind struct {
name string
extensions []string
mimes []string
signatures [][]byte
}
var fileKinds = []fileKind{
{"png", []string{".png"}, []string{"image/png"}, [][]byte{{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}}},
{"jpg", []string{".jpg", ".jpeg"}, []string{"image/jpeg"}, [][]byte{{0xff, 0xd8, 0xff}}},
{"gif", []string{".gif"}, []string{"image/gif"}, [][]byte{[]byte("GIF87a"), []byte("GIF89a")}},
{"webp", []string{".webp"}, []string{"image/webp"}, nil},
{"bmp", []string{".bmp"}, []string{"image/bmp", "image/x-ms-bmp"}, [][]byte{[]byte("BM")}},
{"pdf", []string{".pdf"}, []string{"application/pdf"}, [][]byte{[]byte("%PDF")}},
{"zip", []string{".zip", ".docx", ".xlsx", ".pptx", ".apk", ".jar"},
[]string{"application/zip", "application/x-zip-compressed",
"application/vnd.openxmlformats-officedocument.wordprocessingml.document",
"application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
"application/vnd.openxmlformats-officedocument.presentationml.presentation",
"application/java-archive", "application/vnd.android.package-archive"},
[][]byte{[]byte("PK\x03\x04"), []byte("PK\x05\x06"), []byte("PK\x07\x08")}},
{"rar", []string{".rar"}, []string{"application/x-rar-compressed", "application/vnd.rar"},
[][]byte{[]byte("Rar!\x1a\x07\x00"), []byte("Rar!\x1a\x07\x01\x00")}},
{"7z", []string{".7z"}, []string{"application/x-7z-compressed"}, [][]byte{{0x37, 0x7a, 0xbc, 0xaf, 0x27, 0x1c}}},
{"gz", []string{".gz", ".tgz"}, []string{"application/gzip", "application/x-gzip"}, [][]byte{{0x1f, 0x8b}}},
{"mp3", []string{".mp3"}, []string{"audio/mpeg"}, [][]byte{[]byte("ID3"), {0xff, 0xfb}, {0xff, 0xf3}, {0xff, 0xf2}}},
{"mp4", []string{".mp4", ".m4a", ".mov"}, []string{"video/mp4", "audio/mp4", "video/quicktime"}, nil},
{"exe", []string{".exe", ".dll", ".sys"}, []string{"application/x-msdownload", "application/x-dosexec"}, [][]byte{[]byte("MZ")}},
{"elf", []string{".elf", ".so", ".o"}, []string{"application/x-executable"}, [][]byte{{0x7f, 'E', 'L', 'F'}}},
}
// knownExtensions 全部已知扩展名集合。
var knownExtensions = func() map[string]bool {
m := map[string]bool{}
for _, k := range fileKinds {
for _, ext := range k.extensions {
m[ext] = true
}
}
return m
}()
// isTypeAllowed 判断文件是否在 allowed_file_types 白名单内("*"/*/* 放行全部)。
func isTypeAllowed(cfg *config.Config, fileName, contentType string) bool {
allowed := cfg.AllowedFileTypes()
if len(allowed) == 0 {
return true
}
name := strings.ToLower(strings.TrimSpace(fileName))
ct := strings.ToLower(strings.TrimSpace(contentType))
for _, rule := range allowed {
rule = strings.ToLower(strings.TrimSpace(rule))
switch {
case rule == "*" || rule == "*/*":
return true
case strings.Contains(rule, "/"):
if ok, _ := path.Match(rule, ct); ok {
return true
}
default:
if !strings.HasPrefix(rule, ".") {
rule = "." + rule
}
if strings.HasSuffix(name, rule) {
return true
}
}
}
return false
}
// detectFileKind 按文件头识别类型(对齐参考:RIFF/WEBP、ftyp/mp4 与前缀签名表)。
func detectFileKind(header []byte) *fileKind {
if len(header) == 0 {
return nil
}
if len(header) >= 12 && string(header[:4]) == "RIFF" && string(header[8:12]) == "WEBP" {
for i := range fileKinds {
if fileKinds[i].name == "webp" {
return &fileKinds[i]
}
}
}
if len(header) >= 12 && string(header[4:8]) == "ftyp" {
for i := range fileKinds {
if fileKinds[i].name == "mp4" {
return &fileKinds[i]
}
}
}
var best *fileKind
bestLen := 0
for i := range fileKinds {
for _, sig := range fileKinds[i].signatures {
if len(sig) > 0 && len(header) >= len(sig) && string(header[:len(sig)]) == string(sig) {
if len(sig) > bestLen {
bestLen = len(sig)
best = &fileKinds[i]
}
}
}
}
return best
}
// validateFileMagic 白名单 + magic bytes 防伪造(对齐参考 validate_file_magic)。
// header 为文件前 64 字节,可为空(空则只校验白名单)。
func validateFileMagic(cfg *config.Config, fileName, contentType string, header []byte) error {
if !isTypeAllowed(cfg, fileName, contentType) {
return errForbidden("不允许上传该类型文件")
}
if len(header) == 0 {
return nil
}
ext := strings.ToLower(path.Ext(fileName))
ct := strings.ToLower(strings.TrimSpace(contentType))
detected := detectFileKind(header)
if knownExtensions[ext] {
if detected == nil {
return errForbidden("文件内容与扩展名不匹配,疑似伪造类型")
}
matched := false
for _, e := range detected.extensions {
if e == ext {
matched = true
break
}
}
if !matched {
return errForbidden("文件内容与扩展名不匹配,疑似伪造类型")
}
}
if ct != "" {
for _, k := range fileKinds {
for _, m := range k.mimes {
if m == ct {
// 声明了已知 MIME:内容必须匹配
if detected == nil {
return errForbidden("文件内容与 Content-Type 不匹配,疑似伪造类型")
}
matched := false
for _, m2 := range detected.mimes {
if m2 == ct {
matched = true
break
}
}
if !matched {
return errForbidden("文件内容与 Content-Type 不匹配,疑似伪造类型")
}
break
}
}
}
}
return nil
}
// readMultipartHeader 读取上传文件前 n 字节并 seek 回起点(用于 magic 校验)。
func readMultipartHeader(f multipart.File, n int64) []byte {
if f == nil {
return nil
}
buf := make([]byte, n)
nread, _ := f.Read(buf)
_, _ = f.Seek(0, 0)
if nread <= 0 {
return nil
}
return buf[:nread]
}
// ============ 杂项 ============
// humanSize 把字节数转成人类可读描述(B/KB/MB/GB 自适应;
// v2 需求④:max_file_size 支持子 MB 上限,固定 MB 格式会显示 0.00 MB)。
func humanSize(n int64) string {
const kb, mb, gb = int64(1024), int64(1024 * 1024), int64(1024 * 1024 * 1024)
switch {
case n >= gb:
return fmt.Sprintf("%.2f GB", float64(n)/float64(gb))
case n >= mb:
return fmt.Sprintf("%.2f MB", float64(n)/float64(mb))
case n >= kb:
return fmt.Sprintf("%.2f KB", float64(n)/float64(kb))
default:
return fmt.Sprintf("%d B", n)
}
}
// contentDisposition 生成 RFC 5987 附件头(对齐参考 filename*=UTF-8” 格式)。
func contentDisposition(name string) string {
quoted := urlPathEscape(name)
return "attachment; filename*=UTF-8''" + quoted
}
// urlPathEscape RFC 5987 风格百分号编码(等价 urllib.parse.quote(safe=”))。
func urlPathEscape(s string) string {
var b strings.Builder
for _, r := range []byte(s) {
if (r >= 'A' && r <= 'Z') || (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') ||
r == '-' || r == '_' || r == '.' || r == '~' {
b.WriteByte(r)
} else {
fmt.Fprintf(&b, "%%%02X", r)
}
}
return b.String()
}
// parseISOTime 解析 ISO 8601 / RFC3339 / 日期时间字符串,失败返回错误。
func parseISOTime(s string) (time.Time, error) {
s = strings.TrimSpace(s)
if s == "" {
return time.Time{}, errors.New("空时间")
}
layouts := []string{
time.RFC3339Nano, time.RFC3339,
"2006-01-02T15:04:05", "2006-01-02 15:04:05", "2006-01-02",
}
for _, layout := range layouts {
if t, err := time.Parse(layout, s); err == nil {
return t, nil
}
}
return time.Time{}, fmt.Errorf("时间格式错误: %s", s)
}
// requireUploadLimit 上传限流入口检查(进入即校验,超限 423;成功后由 handler 显式 Add)。
// 对齐参考:FastAPI Depends(ip_limit["upload"]) 在进入时 check。
func requireUploadLimit(c *gin.Context, limiter *middleware.RateLimiter) bool {
if allowed, _ := limiter.Check(c, middleware.LimitUpload); !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return false
}
return true
}
// bindJSONOrForm 兼容 JSON 与表单请求体绑定到结构体。
func bindJSONOrForm(c *gin.Context, obj any) error {
ct := c.GetHeader("Content-Type")
if strings.Contains(ct, "application/json") {
if err := c.ShouldBindJSON(obj); err != nil {
return errBadRequest("请求体格式错误: " + err.Error())
}
return nil
}
// v3.1.1 兼容归一化:旧前端 bundlefetch 字符串 body 默认 text/plain)发的
// 是 text/plain + urlencoded 格式。此类请求改写 Content-Type 后走表单绑定,
// 否则 ShouldBind 对 text/plain 不解析,非空字段全部丢失。
base := ct
if i := strings.IndexByte(ct, ';'); i >= 0 {
base = ct[:i]
}
if strings.EqualFold(strings.TrimSpace(base), "text/plain") &&
c.Request != nil && c.Request.Body != nil {
if raw, err := io.ReadAll(c.Request.Body); err == nil {
trimmed := bytes.TrimSpace(raw)
// JSON 形态(无头/误标 text/plain):改写后按 JSON 绑定(须先于 ParseQuery 判断,
// 否则形如 {"a":1} 的 JSON 会被 ParseQuery 误判为单键 urlencoded
if len(trimmed) > 0 && trimmed[0] == '{' {
c.Request.Header.Set("Content-Type", "application/json")
c.Request.Body = io.NopCloser(bytes.NewReader(raw))
return c.ShouldBindJSON(obj)
}
if vals, perr := url.ParseQuery(string(raw)); perr == nil && len(vals) > 0 {
// urlencoded 形态:改写 Content-Type 走表单绑定
c.Request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
c.Request.Body = io.NopCloser(bytes.NewReader(raw))
return c.ShouldBind(obj)
}
// 其他形态:还原 body 让 ShouldBind 按原样处理
c.Request.Body = io.NopCloser(bytes.NewReader(raw))
}
}
if err := c.ShouldBind(obj); err != nil {
return errBadRequest("请求体格式错误: " + err.Error())
}
return nil
}
// auditUploadEntry 填充上传类审计业务字段的便捷函数。
func auditUploadEntry(c *gin.Context, code, name string, size, transferred int64) {
middleware.AuditSet(c, func(e *audit.Entry) {
e.FileCode = code
e.FileName = name
e.SizeBytes = size
e.TransferredBytes = transferred
})
}
// auditRecordSuccess / auditRecordFailed 显式落库便捷函数。
func auditRecordSuccess(c *gin.Context, svc *audit.Service) {
middleware.AuditRecordRequest(c, svc, model.AuditResultSuccess, "")
}
func auditRecordFailed(c *gin.Context, svc *audit.Service, msg string) {
middleware.AuditRecordRequest(c, svc, model.AuditResultFailed, msg)
}
// uuidHex 生成 32 位十六进制随机串(对齐参考 uuid4().hex)。
func uuidHex() string {
b := make([]byte, 16)
_, _ = rand.Read(b)
// 设置版本号与变体位以保持 uuid4 兼容格式
b[6] = (b[6] & 0x0f) | 0x40
b[8] = (b[8] & 0x3f) | 0x80
return hex.EncodeToString(b)
}
// uuidCanonical 生成带连字符的 UUID 字符串(upload_id 用)。
func uuidCanonical() string {
h := uuidHex()
return h[0:8] + "-" + h[8:12] + "-" + h[12:16] + "-" + h[16:20] + "-" + h[20:32]
}
+264
View File
@@ -0,0 +1,264 @@
package api
import (
"errors"
"fmt"
"testing"
"time"
"filecodebox/internal/config"
"filecodebox/internal/storage"
)
// newTestConfig 构造测试配置(defaults 基线,无 KV 覆盖;需求 ⑧ 默认 sqlite,无需真实数据库)。
func newTestConfig(t *testing.T) *config.Config {
t.Helper()
t.Setenv("FCB_DB_DRIVER", "sqlite")
t.Setenv("FCB_DB_DSN", "")
cfg, err := config.New()
if err != nil {
t.Fatalf("config.New 失败: %v", err)
}
return cfg
}
// TestGetSelectTokenWindow 验证下载令牌的窗口号逻辑(对齐参考 get_select_token)。
func TestGetSelectTokenWindow(t *testing.T) {
code := "AB12C"
secret := "test-secret"
tok0 := GetSelectToken(code, secret, 0)
tok1 := GetSelectToken(code, secret, 1)
if tok0 == "" || tok1 == "" {
t.Fatal("令牌不应为空")
}
// 同一窗口内 offset=0 的两次生成必须一致(确定性问题)
if tok0 != GetSelectToken(code, secret, 0) {
t.Fatal("同窗口令牌应确定一致")
}
// 不同 offset 的令牌必然不同(窗口号不同)
if tok0 == tok1 {
t.Fatal("offset=0 与 offset=1 的令牌应不同")
}
// 不同 code / secret 的令牌不同
if tok0 == GetSelectToken("ZZ999", secret, 0) {
t.Fatal("不同取件码的令牌应不同")
}
if tok0 == GetSelectToken(code, "other-secret", 0) {
t.Fatal("不同密钥的令牌应不同")
}
// 令牌为 64 位十六进制(sha256 hex
if len(tok0) != 64 {
t.Fatalf("令牌长度应为 64,实际 %d", len(tok0))
}
// 窗口号公式:unix/1000 - offset(秒级窗口约 16.7 分钟)
now := time.Now().Unix()
if now/1000 == (now+1100)/1000 {
// 仅当测试跨越窗口边界才跳过该断言(罕见,容忍)
t.Log("测试跨越窗口边界,跳过窗口公式断言")
}
}
// TestMapStorageError 验证存储哨兵错误→HTTP 状态映射表。
func TestMapStorageError(t *testing.T) {
cases := []struct {
name string
err error
expect int
}{
{"NotFound 直返", storage.ErrNotFound, 404},
{"NotFound 包装", fmt.Errorf("引擎内层: %w", storage.ErrNotFound), 404},
{"InvalidPath", storage.ErrInvalidPath, 400},
{"InvalidPath 包装", fmt.Errorf("webdav: %w", storage.ErrInvalidPath), 400},
{"Unavailable", storage.ErrUnavailable, 503},
{"Unavailable 包装", fmt.Errorf("s3: %w", storage.ErrUnavailable), 503},
{"NotSupported", storage.ErrNotSupported, 501},
{"NotSupported 包装", fmt.Errorf("local: %w", storage.ErrNotSupported), 501},
{"RangeNotSatisfiable", storage.ErrRangeNotSatisfiable, 416},
{"HashMismatch", storage.ErrHashMismatch, 400},
{"未知错误", errors.New("其他错误"), 500},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
mapped := mapStorageError(tc.err)
var ae *apiError
if !errors.As(mapped, &ae) {
t.Fatalf("应映射为 apiError,得到 %T", mapped)
}
if ae.Status != tc.expect {
t.Fatalf("状态码应为 %d,实际 %d", tc.expect, ae.Status)
}
})
}
}
// TestResolveExpire 验证过期策略解析(白名单/上限/各 style)。
func TestResolveExpire(t *testing.T) {
cfg := newTestConfig(t)
// 非法 style
if _, err := resolveExpire(cfg, 1, "year"); err == nil {
t.Fatal("非白名单 style 应报错")
}
// 非法 value
if _, err := resolveExpire(cfg, 0, "day"); err == nil {
t.Fatal("expire_value<=0 应报错")
}
// day7 天内合法
res, err := resolveExpire(cfg, 3, "day")
if err != nil {
t.Fatalf("3 天应合法: %v", err)
}
if res.ExpiredAt == nil || res.ExpiredCount != -1 {
t.Fatal("day 类型应有 expired_at 且 expired_count=-1")
}
// 超过 7 天上限
if _, err := resolveExpire(cfg, 30, "day"); err == nil {
t.Fatal("超过 7 天上限应报错")
}
// count:按次数
res, err = resolveExpire(cfg, 5, "count")
if err != nil {
t.Fatalf("count 应合法: %v", err)
}
if res.ExpiredCount != 5 {
t.Fatalf("count 类型 expired_count 应为 5,实际 %d", res.ExpiredCount)
}
// forever:永久
res, err = resolveExpire(cfg, 1, "forever")
if err != nil {
t.Fatalf("forever 应合法: %v", err)
}
if res.ExpiredAt != nil || res.ExpiredCount != -1 {
t.Fatal("forever 应为 expired_at=nil 且 expired_count=-1")
}
}
// TestGenerateCode 验证取件码格式。
func TestGenerateCode(t *testing.T) {
for i := 0; i < 50; i++ {
num := generateCode("number")
if len(num) != 5 {
t.Fatalf("数字码应为 5 位,实际 %q", num)
}
for _, ch := range num {
if ch < '0' || ch > '9' {
t.Fatalf("数字码含非数字字符: %q", num)
}
}
secret := generateCode("secret")
if len(secret) != 5 {
t.Fatalf("字符码应为 5 位,实际 %q", secret)
}
}
}
// TestParseRangeHeader 验证 Range 头解析(对齐 HTTP 语义)。
func TestParseRangeHeader(t *testing.T) {
// 全量(无 Range
if parseRangeHeader("", 1000) != nil {
t.Fatal("无 Range 头应返回 nil")
}
// 标准区间
r := parseRangeHeader("bytes=0-99", 1000)
if r == nil || r.Start != 0 || r.End != 99 {
t.Fatalf("bytes=0-99 解析错误: %+v", r)
}
// 开区间到末尾
r = parseRangeHeader("bytes=500-", 1000)
if r == nil || r.Start != 500 || r.End != -1 {
t.Fatalf("bytes=500- 解析错误: %+v", r)
}
// 后缀区间(最后 100 字节)
r = parseRangeHeader("bytes=-100", 1000)
if r == nil || r.Start != 900 || r.End != -1 {
t.Fatalf("bytes=-100 解析错误: %+v", r)
}
// 后缀超长:截断到全文件
r = parseRangeHeader("bytes=-5000", 1000)
if r == nil || r.Start != 0 {
t.Fatalf("bytes=-5000 应从头开始: %+v", r)
}
// 多区间不支持→回退全量
if parseRangeHeader("bytes=0-1,5-6", 1000) != nil {
t.Fatal("多区间应返回 nil(回退全量)")
}
// 非法格式
if parseRangeHeader("items=0-1", 1000) != nil {
t.Fatal("非 bytes 单位应返回 nil")
}
if parseRangeHeader("bytes=abc-", 1000) != nil {
t.Fatal("非法数字应返回 nil")
}
}
// TestParseISOTime 验证时间解析的多格式兼容。
func TestParseISOTime(t *testing.T) {
valid := []string{
"2025-01-01T00:00:00Z",
"2025-01-01T08:00:00+08:00",
"2025-01-01 08:00:00",
"2025-01-01",
}
for _, s := range valid {
if _, err := parseISOTime(s); err != nil {
t.Fatalf("%q 应解析成功: %v", s, err)
}
}
if _, err := parseISOTime("not-a-time"); err == nil {
t.Fatal("非法时间应报错")
}
if _, err := parseISOTime(""); err == nil {
t.Fatal("空串应报错")
}
}
// TestFormatDurationCN 验证中文时长描述。
func TestFormatDurationCN(t *testing.T) {
cases := []struct {
d time.Duration
expect string
}{
{7 * 24 * time.Hour, "7天"},
{90 * time.Minute, "1小时30分钟"},
{45 * time.Second, "45秒"},
}
for _, tc := range cases {
if got := formatDurationCN(tc.d); got != tc.expect {
t.Fatalf("formatDurationCN(%v)=%q,期望 %q", tc.d, got, tc.expect)
}
}
}
// TestFileMagicValidation 验证 magic bytes 防伪造。
func TestFileMagicValidation(t *testing.T) {
cfg := newTestConfig(t)
// 白名单 * 全放行
if err := validateFileMagic(cfg, "a.txt", "", nil); err != nil {
t.Fatalf("白名单 * 应放行: %v", err)
}
// PNG 内容 + .exe 扩展名 → 拒绝(伪造)
if err := validateFileMagic(cfg, "evil.exe", "", []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}); err == nil {
t.Fatal("PNG 内容伪装 exe 应拒绝")
}
// PNG 内容 + .png 扩展名 → 通过
if err := validateFileMagic(cfg, "ok.png", "image/png", []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}); err != nil {
t.Fatalf("真 PNG 应通过: %v", err)
}
// 文本内容 + .png 扩展名 → 拒绝
if err := validateFileMagic(cfg, "fake.png", "", []byte("hello world, this is text")); err == nil {
t.Fatal("文本伪装 png 应拒绝")
}
}
// TestSanitizePathBuild 验证存储路径构造不含穿越。
func TestSanitizePathBuild(t *testing.T) {
cfg := newTestConfig(t)
_, _, _, clean, savePath := buildSavePath(cfg, "../../etc/passwd", "uuid-123")
if clean != "etc_passwd" && clean != "passwd" {
t.Logf("清理后的文件名: %q", clean)
}
if _, ok := storage.SanitizePath(savePath); !ok {
t.Fatalf("构造的 savePath 应通过安全校验: %q", savePath)
}
}
+45
View File
@@ -0,0 +1,45 @@
// policy.go — v2 上传策略统一读取与校验(需求 ④⑩)。
//
// 管理端在后台设置页修改策略(settings KVt1 schema)后,上传链路
// share/file、chunk、presign)每次请求实时读取当前策略并动态校验:
// - 大小上限:max_file_size0=回落 uploadSize,语义见 config.MaxFileSize);
// - 类型白名单:allowed_file_types"*" 不限制),由 validateFileMagic 统一执行;
// - 保存策略:expire_style 白名单 + max_save_seconds 时间上限 + max_save_count
// 次数上限,统一在 resolveExpirehelpers.go)执行。
//
// 超限返回 403(超出策略限制)/400(参数非法),错误信息为中文。
package api
import (
"fmt"
)
// UploadPolicy 当前生效的上传策略快照(每次上传请求实时读取,管理端改动立即生效)。
type UploadPolicy struct {
MaxFileSize int64 // 单文件大小上限(字节),0=不限制
AllowedTypes []string // 类型白名单,"*" 不限制
ExpireStyles []string // 允许的过期方式白名单
MaxSaveSeconds int64 // 最长保存秒数,0=不限制(默认 7 天兜底)
MaxSaveCount int // 单次分享最大可取次数上限,0=不限制
}
// CurrentUploadPolicy 读取当前上传策略快照。
// 上传页亦通过 GET /api/v1/config 的 policy 字段读取同一组值做动态渲染。
func (d *Deps) CurrentUploadPolicy() UploadPolicy {
cfg := d.Cfg
return UploadPolicy{
MaxFileSize: cfg.MaxFileSize(),
AllowedTypes: cfg.AllowedFileTypes(),
ExpireStyles: cfg.ExpireStyle(),
MaxSaveSeconds: cfg.MaxSaveSeconds(),
MaxSaveCount: cfg.MaxSaveCount(),
}
}
// CheckSize 校验单文件大小是否超出策略上限(超出返回 403,文案对齐参考实现)。
func (p UploadPolicy) CheckSize(size int64) error {
if p.MaxFileSize > 0 && size > p.MaxFileSize {
return errForbidden(fmt.Sprintf("大小超过限制,最大为%s", humanSize(p.MaxFileSize)))
}
return nil
}
+571
View File
@@ -0,0 +1,571 @@
package api
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"mime/multipart"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/cache"
"filecodebox/internal/config"
"filecodebox/internal/database"
"filecodebox/internal/middleware"
"filecodebox/internal/settings"
"filecodebox/internal/storage"
)
// ============ 测试环境装配(真实 sqlite + 内存缓存 + 本地存储)============
// newPolicyTestDeps 构造带真实依赖的 Depssqlite 文件库(t.TempDir)、
// 本地存储引擎、内存缓存限流器与审计服务(需求 ⑧ 默认形态)。
func newPolicyTestDeps(t *testing.T) *Deps {
t.Helper()
gin.SetMode(gin.TestMode)
dir := t.TempDir()
t.Setenv("FCB_DB_DRIVER", "sqlite")
t.Setenv("FCB_DB_DSN", filepath.Join(dir, "test.db"))
cfg, err := config.New()
if err != nil {
t.Fatalf("config.New: %v", err)
}
ctx := context.Background()
db, err := database.Open(ctx, database.Options{Driver: config.DBDriverSQLite, DSN: filepath.Join(dir, "test.db")})
if err != nil {
t.Fatalf("database.Open: %v", err)
}
t.Cleanup(func() { _ = database.Close(db) })
if err := database.Migrate(ctx, db); err != nil {
t.Fatalf("database.Migrate: %v", err)
}
mgr, err := settings.NewManager(ctx, db, cfg)
if err != nil {
t.Fatalf("settings.NewManager: %v", err)
}
store, err := storage.NewLocalStorage(filepath.Join(dir, "storage"))
if err != nil {
t.Fatalf("storage.NewLocalStorage: %v", err)
}
// v3:包装为 Managerbuild 直接返回 local 实例,测试无需真实多引擎)
storeMgr := storage.NewManager("local", store, func(string) (storage.Storage, error) {
return storage.NewLocalStorage(filepath.Join(dir, "storage"))
})
return &Deps{
DB: db,
Cfg: cfg,
Mgr: mgr,
AuditSvc: audit.NewService(audit.NewDBSink(db)),
Limiter: middleware.NewRateLimiter(cache.NewMemory(), nil),
Store: storeMgr,
Version: "test",
}
}
// ============ 请求构造辅助 ============
// invoke 以给定请求调用 handler 并返回响应。
func invoke(handler gin.HandlerFunc, req *http.Request) *httptest.ResponseRecorder {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
handler(c)
return w
}
// patchConfig 以 JSON 调用 PATCH /admin/config/update。
func patchConfig(d *Deps, patch map[string]any) *httptest.ResponseRecorder {
raw, _ := json.Marshal(patch)
req := httptest.NewRequest(http.MethodPatch, "/admin/config/update", bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
return invoke(d.adminConfigUpdate, req)
}
// getConfig 调用 GET /admin/config/get。
func getConfig(d *Deps) *httptest.ResponseRecorder {
return invoke(d.adminConfigGet, httptest.NewRequest(http.MethodGet, "/admin/config/get", nil))
}
// getPublicConfig 调用 GET /api/v1/config。
func getPublicConfig(d *Deps) *httptest.ResponseRecorder {
return invoke(d.publicConfig, httptest.NewRequest(http.MethodGet, "/api/v1/config", nil))
}
// uploadFile 以 multipart 表单调用 POST /share/file。
func uploadFile(d *Deps, name string, content []byte, fields map[string]string) *httptest.ResponseRecorder {
body := &bytes.Buffer{}
mw := multipart.NewWriter(body)
fw, err := mw.CreateFormFile("file", name)
if err != nil {
panic(err)
}
_, _ = fw.Write(content)
for k, v := range fields {
_ = mw.WriteField(k, v)
}
_ = mw.Close()
req := httptest.NewRequest(http.MethodPost, "/share/file", body)
req.Header.Set("Content-Type", mw.FormDataContentType())
return invoke(d.shareFile, req)
}
// chunkInitJSON 以 JSON 调用 POST /chunk/upload/init。
func chunkInitJSON(d *Deps, payload string) *httptest.ResponseRecorder {
req := httptest.NewRequest(http.MethodPost, "/chunk/upload/init", strings.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
return invoke(d.chunkInit, req)
}
// respBody 解析统一响应体。
func respBody(t *testing.T, w *httptest.ResponseRecorder) (code int, data map[string]any) {
t.Helper()
var body struct {
Code int `json:"code"`
Msg string `json:"msg"`
Data map[string]any `json:"data"`
}
if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil {
t.Fatalf("响应解析失败: %v; body=%s", err, w.Body.String())
}
return body.Code, body.Data
}
// pngMagic 最小合法 PNG 头(magic 校验可识别)。
var pngMagic = []byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'}
// ============ ① 公开 configv2 展示与策略字段下发 ============
// TestPublicConfigV2Fields 验证 /api/v1/config 下发背景/页脚/备案/通知与策略范围,
// 且响应不包含任何敏感键(admin_token/jwt_secret)。
func TestPublicConfigV2Fields(t *testing.T) {
d := newPolicyTestDeps(t)
// 管理端先设置 v2 展示字段
if w := patchConfig(d, map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "自定义页脚内容",
"footer_beian": "京ICP备2024xxxxxx号-1",
"notify_enabled": 0,
"max_save_count": 5,
}); w.Code != 200 {
t.Fatalf("patchConfig 失败: %d %s", w.Code, w.Body.String())
}
w := getPublicConfig(d)
code, data := respBody(t, w)
if code != 200 {
t.Fatalf("publicConfig code=%d", code)
}
cfgMap, _ := data["config"].(map[string]any)
if cfgMap == nil {
t.Fatal("响应缺少 config 对象")
}
for key, want := range map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "自定义页脚内容",
"footer_beian": "京ICP备2024xxxxxx号-1",
"notify_enabled": float64(0),
"notify_title": "系统通知",
} {
if got := cfgMap[key]; got != want {
t.Fatalf("config.%s = %v, 期望 %v", key, got, want)
}
}
// 策略范围
if _, ok := cfgMap["max_file_size"]; !ok {
t.Fatal("config 缺少 max_file_size(存储策略)")
}
if _, ok := cfgMap["max_save_seconds"]; !ok {
t.Fatal("config 缺少 max_save_seconds(保存时间策略)")
}
if got := cfgMap["max_save_count"]; got != float64(5) {
t.Fatalf("config.max_save_count = %v, 期望 5", got)
}
if _, ok := cfgMap["allowedFileTypes"]; !ok {
t.Fatal("config 缺少 allowedFileTypes")
}
if _, ok := cfgMap["expireStyle"]; !ok {
t.Fatal("config 缺少 expireStyle")
}
if _, ok := cfgMap["uploadSize"]; !ok {
t.Fatal("config 缺少 uploadSize")
}
// 敏感键绝不下发
raw := w.Body.String()
if strings.Contains(raw, "admin_token") || strings.Contains(raw, "jwt_secret") {
t.Fatal("公开 config 响应包含敏感键")
}
}
// ============ ② 管理端 get/updatev2 键全链路 + 类型范围校验 ============
// TestAdminConfigV2RoundTrip 验证 v2 新键 update → get → public 的往返,
// 且 admin_token 屏蔽、jwt_secret 不下发。
func TestAdminConfigV2RoundTrip(t *testing.T) {
d := newPolicyTestDeps(t)
patch := map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "页脚 HTML 片段",
"footer_beian": "京ICP备20240001号",
"notify_enabled": 0,
"max_save_count": 20,
"max_file_size": 5242880,
"max_save_seconds": 86400,
}
w := patchConfig(d, patch)
if w.Code != 200 {
t.Fatalf("update 失败: %d %s", w.Code, w.Body.String())
}
// admin config get:新键可见 + 敏感键屏蔽
w = getConfig(d)
code, data := respBody(t, w)
if code != 200 {
t.Fatalf("get code=%d", code)
}
for key, want := range map[string]any{
"background_url": "https://cdn.example.com/bg.png",
"footer_text": "页脚 HTML 片段",
"footer_beian": "京ICP备20240001号",
"notify_enabled": float64(0),
"max_save_count": float64(20),
"max_file_size": float64(5242880),
"max_save_seconds": float64(86400),
} {
if got := data[key]; got != want {
t.Fatalf("admin get %s = %v, 期望 %v", key, got, want)
}
}
// v1 既有设计:admin_token 不在 configKeys 白名单(响应中不存在即屏蔽);
// 兼容两种形态:键缺失或空串均算通过
if got, present := data["admin_token"]; present && got != "" {
t.Fatalf("admin_token 应屏蔽(缺失或空串),实际 %v", got)
}
rawGet := w.Body.String()
if strings.Contains(rawGet, `"jwt_secret"`) {
t.Fatal("admin get 不应下发 jwt_secret")
}
// public config 立即反映(改策略 → 公开 config 即时更新)
w = getPublicConfig(d)
_, data = respBody(t, w)
cfgMap := data["config"].(map[string]any)
if got := cfgMap["max_file_size"]; got != float64(5242880) {
t.Fatalf("public max_file_size = %v, 期望 5242880", got)
}
if got := cfgMap["footer_beian"]; got != "京ICP备20240001号" {
t.Fatalf("public footer_beian = %v", got)
}
}
// TestAdminConfigV2Validation 验证新键的类型与范围校验(400 + 中文错误)。
func TestAdminConfigV2Validation(t *testing.T) {
d := newPolicyTestDeps(t)
cases := []struct {
name string
patch map[string]any
}{
{"max_file_size 负数", map[string]any{"max_file_size": -1}},
{"max_file_size 超上限", map[string]any{"max_file_size": config.MaxFileSizeMax + 1}},
{"max_save_count 超上限", map[string]any{"max_save_count": config.MaxSaveCountMax + 1}},
{"notify_enabled 越界", map[string]any{"notify_enabled": 2}},
{"footer_beian 超长", map[string]any{"footer_beian": strings.Repeat("备", config.FooterBeianMaxLen+1)}},
{"footer_text 超长", map[string]any{"footer_text": strings.Repeat("页", config.FooterTextMaxLen+1)}},
{"background_url 非法协议", map[string]any{"background_url": "javascript:alert(1)"}},
{"max_save_seconds 超上限", map[string]any{"max_save_seconds": config.MaxSaveSecondsMax + 1}},
{"allowed_file_types 类型错误", map[string]any{"allowed_file_types": 123}},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
w := patchConfig(d, tc.patch)
if w.Code != http.StatusBadRequest {
t.Fatalf("应 400,实际 %d %s", w.Code, w.Body.String())
}
code, _ := respBody(t, w)
if code != http.StatusBadRequest {
t.Fatalf("响应 code 应为 400,实际 %d", code)
}
})
}
// 合法值不受影响
if w := patchConfig(d, map[string]any{
"max_save_count": 0, // 0 = 不限制
"notify_enabled": 1,
"background_url": "data:image/png;base64,AAA",
"max_file_size": 1024,
"max_save_seconds": 0,
}); w.Code != 200 {
t.Fatalf("合法 patch 应 200: %d %s", w.Code, w.Body.String())
}
}
// ============ ③ 上传动态校验:admin 改策略 → 上传行为即时变化 ============
// TestUploadPolicyDynamicEnforcement 全链路:默认可传 → 改 max_file_size/白名单/
// 保存策略后 → 公开 config 反映 → 上传被新策略拒绝(403/400)。
func TestUploadPolicyDynamicEnforcement(t *testing.T) {
d := newPolicyTestDeps(t)
// 默认策略:小 PNG 上传成功
w := uploadFile(d, "ok.png", pngMagic, map[string]string{"expire_value": "1", "expire_style": "day"})
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("默认策略上传应成功: %d %s", w.Code, w.Body.String())
}
// —— 大小上限:max_file_size=100 → 200B 文件 403 ——
if w = patchConfig(d, map[string]any{"max_file_size": 100}); w.Code != 200 {
t.Fatalf("patch max_file_size: %d %s", w.Code, w.Body.String())
}
w = getPublicConfig(d)
_, data := respBody(t, w)
if got := data["config"].(map[string]any)["max_file_size"]; got != float64(100) {
t.Fatalf("公开 config 未即时反映 max_file_size=100: %v", got)
}
w = uploadFile(d, "big.png", append(pngMagic, bytes.Repeat([]byte{0}, 200)...),
map[string]string{"expire_value": "1", "expire_style": "day"})
code, _ := respBody(t, w)
if code != http.StatusForbidden {
t.Fatalf("超限上传应 403: %d %s", w.Code, w.Body.String())
}
// —— 类型白名单:allowed_file_types=[.png] → .txt 403 ——
if w = patchConfig(d, map[string]any{"allowed_file_types": []string{".png"}}); w.Code != 200 {
t.Fatalf("patch allowed_file_types: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "note.txt", []byte("hello"), map[string]string{"expire_value": "1", "expire_style": "day"})
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("非白名单类型应 403: %d %s", w.Code, w.Body.String())
}
// —— 保存时间:max_save_seconds=3600 → expire 1 天 403 ——
if w = patchConfig(d, map[string]any{"max_save_seconds": 3600}); w.Code != 200 {
t.Fatalf("patch max_save_seconds: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "timed.png", pngMagic, map[string]string{"expire_value": "1", "expire_style": "day"})
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("保存时间超范围应 403: %d %s", w.Code, w.Body.String())
}
// —— 保存次数:max_save_count=5 → count=10 403(重置时间策略避免交叉影响)——
if w = patchConfig(d, map[string]any{"max_save_count": 5, "max_save_seconds": 0}); w.Code != 200 {
t.Fatalf("patch max_save_count: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "counted.png", pngMagic, map[string]string{"expire_value": "10", "expire_style": "count"})
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("保存次数超上限应 403: %d %s", w.Code, w.Body.String())
}
// 次数在上限内合法
w = uploadFile(d, "counted.png", pngMagic, map[string]string{"expire_value": "3", "expire_style": "count"})
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("次数在上限内应成功: %d %s", w.Code, w.Body.String())
}
// —— 过期方式白名单收窄:expireStyle=[day] → hour 400 ——
if w = patchConfig(d, map[string]any{"expireStyle": []string{"day"}}); w.Code != 200 {
t.Fatalf("patch expireStyle: %d %s", w.Code, w.Body.String())
}
w = uploadFile(d, "hour.png", pngMagic, map[string]string{"expire_value": "2", "expire_style": "hour"})
if code, _ := respBody(t, w); code != http.StatusBadRequest {
t.Fatalf("非白名单 expire_style 应 400: %d %s", w.Code, w.Body.String())
}
// 恢复后可用(说明策略动态读取)
if w = patchConfig(d, map[string]any{"expireStyle": []string{"day", "hour", "minute", "forever", "count"}}); w.Code != 200 {
t.Fatal("恢复 expireStyle 失败")
}
w = uploadFile(d, "hour.png", pngMagic, map[string]string{"expire_value": "2", "expire_style": "hour"})
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("白名单恢复后应成功: %d %s", w.Code, w.Body.String())
}
}
// TestChunkUploadPolicyEnforcement 验证分片上传链路接入动态策略。
func TestChunkUploadPolicyEnforcement(t *testing.T) {
d := newPolicyTestDeps(t)
// L4:后端强制 enableChunk 开关,本测试前置开启
if w := patchConfig(d, map[string]any{"enableChunk": 1}); w.Code != 200 {
t.Fatalf("patch enableChunk: %d %s", w.Code, w.Body.String())
}
// 大小:max_file_size=1000 → file_size 5000 拒绝
if w := patchConfig(d, map[string]any{"max_file_size": 1000}); w.Code != 200 {
t.Fatalf("patch max_file_size: %d %s", w.Code, w.Body.String())
}
w := chunkInitJSON(d, `{"file_name":"a.png","file_size":5000,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("分片总大小超限应 403: %d %s", w.Code, w.Body.String())
}
// 类型:allowed_file_types=[.png] → b.txt 拒绝(max_file_size 重置为回落,隔离类型断言)
if w := patchConfig(d, map[string]any{"allowed_file_types": []string{".png"}, "max_file_size": 0}); w.Code != 200 {
t.Fatalf("patch allowed_file_types: %d %s", w.Code, w.Body.String())
}
w = chunkInitJSON(d, `{"file_name":"b.txt","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("分片文件类型非白名单应 403: %d %s", w.Code, w.Body.String())
}
// 白名单内 + 大小内 → 会话创建成功
w = chunkInitJSON(d, `{"file_name":"c.png","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("合法分片初始化应 200: %d %s", w.Code, w.Body.String())
}
}
// TestPresignPolicyEnforcement 验证预签名直传链路接入动态大小策略。
func TestPresignPolicyEnforcement(t *testing.T) {
d := newPolicyTestDeps(t)
if w := patchConfig(d, map[string]any{"max_file_size": 1000}); w.Code != 200 {
t.Fatalf("patch max_file_size: %d %s", w.Code, w.Body.String())
}
payload := `{"file_name":"a.png","file_size":5000,"expire_value":1,"expire_style":"day"}`
req := httptest.NewRequest(http.MethodPost, "/presign/upload/init", strings.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
w := invoke(d.presignInit, req)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("预签名直传超限应 403: %d %s", w.Code, w.Body.String())
}
}
// TestSensitiveKeysNeverInPublicConfig 额外兜底:公开 config 任意策略下都无敏感键。
func TestSensitiveKeysNeverInPublicConfig(t *testing.T) {
d := newPolicyTestDeps(t)
// 写入敏感 KV(模拟已初始化实例),公开端点依旧不能带出
ctx := context.Background()
if err := d.Mgr.UpdateKV(ctx, map[string]any{
"jwt_secret": "super-secret-value",
"admin_token": settings.HashPassword("password-123456"),
}); err != nil {
t.Fatalf("UpdateKV: %v", err)
}
if err := d.Mgr.Reload(ctx); err != nil {
t.Fatalf("Reload: %v", err)
}
raw := getPublicConfig(d).Body.String()
if strings.Contains(raw, "super-secret-value") || strings.Contains(raw, "jwt_secret") {
t.Fatal("公开 config 泄露 jwt_secret")
}
if strings.Contains(raw, "admin_token") {
t.Fatal("公开 config 泄露 admin_token")
}
}
// TestPolicySnapshotMatchesConfig 验证策略快照与 config 一致(单一读取口径)。
func TestPolicySnapshotMatchesConfig(t *testing.T) {
d := newPolicyTestDeps(t)
if w := patchConfig(d, map[string]any{"max_file_size": 2048, "max_save_count": 9, "max_save_seconds": 7200}); w.Code != 200 {
t.Fatal("patch 失败")
}
pol := d.CurrentUploadPolicy()
if pol.MaxFileSize != 2048 || pol.MaxSaveCount != 9 || pol.MaxSaveSeconds != 7200 {
t.Fatalf("策略快照不一致: %+v", pol)
}
if err := pol.CheckSize(2048); err != nil {
t.Fatalf("边界值应放行: %v", err)
}
if err := pol.CheckSize(2049); err == nil {
t.Fatal("超限应拒绝")
}
// 0=回落 uploadSize
if w := patchConfig(d, map[string]any{"max_file_size": 0}); w.Code != 200 {
t.Fatal("patch 失败")
}
if got := d.CurrentUploadPolicy().MaxFileSize; got != d.Cfg.UploadSize() {
t.Fatalf("max_file_size=0 应回落 uploadSize: %d vs %d", got, d.Cfg.UploadSize())
}
}
// 编译期保证 fmt 被使用(测试辅助函数中错误路径占位)。
var _ = fmt.Sprintf
// ============ v3 存储引擎热切换 ============
// switchEngine 调用 POST /admin/storage/switch。
func switchEngine(d *Deps, engine string) *httptest.ResponseRecorder {
raw, _ := json.Marshal(map[string]any{"engine": engine})
req := httptest.NewRequest(http.MethodPost, "/admin/storage/switch", bytes.NewReader(raw))
req.Header.Set("Content-Type", "application/json")
return invoke(d.adminStorageSwitch, req)
}
// TestAdminStorageSwitchLocal 本地引擎切换(测试 build 只产 local,切 local 恒成功)。
func TestAdminStorageSwitchLocal(t *testing.T) {
d := newPolicyTestDeps(t)
w := switchEngine(d, "local")
if w.Code != 200 {
t.Fatalf("switch local 失败: %d %s", w.Code, w.Body.String())
}
// KV 持久化:admin get 可见
if got := getConfig(d); got.Code != 200 {
t.Fatal("get 失败")
}
// 非法引擎名 400
if w := switchEngine(d, "ftp"); w.Code != http.StatusBadRequest {
t.Fatalf("非法引擎应 400,实际 %d", w.Code)
}
}
// TestAdminConfigEngineSwitchFailure 测试环境下切换到不可用引擎保持原引擎(503)。
// 测试 Manager 的 build 返回 local;这里通过直接操作 Manager 验证 503 路径的响应格式。
func TestAdminConfigEngineSwitchFailure(t *testing.T) {
d := newPolicyTestDeps(t)
// 用一个恒失败的 Manager 替换(模拟 s3/webdav 健康检查不过)
d.Store = storage.NewManager("local", mustLocal(t), func(string) (storage.Storage, error) {
return nil, errors.New("连接失败")
})
w := switchEngine(d, "s3")
if w.Code != http.StatusServiceUnavailable {
t.Fatalf("不可用引擎应 503,实际 %d %s", w.Code, w.Body.String())
}
if !strings.Contains(w.Body.String(), "已保持原引擎") {
t.Fatal("错误信息应包含「已保持原引擎」")
}
// 失败后当前引擎不变
if d.Store.CurrentName() != "local" {
t.Fatalf("失败后应保持 local,实际 %s", d.Store.CurrentName())
}
}
// TestAdminConfigMaskedSecrets 敏感引擎凭据:get 掩码、update 空/掩码不落库。
func TestAdminConfigMaskedSecrets(t *testing.T) {
d := newPolicyTestDeps(t)
// 先写入真实凭据
if w := patchConfig(d, map[string]any{"webdav_password": "real-secret", "s3_secret_access_key": "sk-real"}); w.Code != 200 {
t.Fatalf("写凭据失败: %s", w.Body.String())
}
// get 应为掩码
_, data := respBody(t, getConfig(d))
if got := data["webdav_password"]; got != settings.SensitiveMaskValue {
t.Fatalf("webdav_password 应掩码,实际 %v", got)
}
if got := data["s3_secret_access_key"]; got != settings.SensitiveMaskValue {
t.Fatalf("s3_secret_access_key 应掩码,实际 %v", got)
}
// 提交掩码(模拟前端回显原样提交)→ 不应覆盖为掩码串
if w := patchConfig(d, map[string]any{"webdav_password": settings.SensitiveMaskValue}); w.Code != 200 {
t.Fatalf("掩码提交应 200: %s", w.Body.String())
}
// 提交空串 → 不修改
if w := patchConfig(d, map[string]any{"s3_secret_access_key": ""}); w.Code != 200 {
t.Fatalf("空串提交应 200: %s", w.Body.String())
}
// 公开 config 绝不含引擎凭据
_, pub := respBody(t, getPublicConfig(d))
rawPub, _ := json.Marshal(pub)
for _, sk := range []string{"webdav_password", "s3_secret_access_key", "aws_session_token", "jwt_secret"} {
if strings.Contains(string(rawPub), sk) {
t.Fatalf("公开 config 不应包含 %s", sk)
}
}
}
func mustLocal(t *testing.T) storage.Storage {
t.Helper()
s, err := storage.NewLocalStorage(t.TempDir())
if err != nil {
t.Fatal(err)
}
return s
}
+511
View File
@@ -0,0 +1,511 @@
package api
import (
"errors"
"fmt"
"net/http"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// presignSessionExpires 预签名会话有效期(对齐参考 PRESIGN_SESSION_EXPIRES=900 秒)。
const presignSessionExpires = 900
// getValidSession 校验并返回预签名会话(对齐参考 _get_valid_session):
// 不存在 404、已过期删除后 404、mode 不符 400。
func (d *Deps) getValidSession(c *gin.Context, uploadID, expectedMode string) (*model.PresignUploadSession, error) {
ctx := c.Request.Context()
var session model.PresignUploadSession
err := d.DB.WithContext(ctx).
Where("upload_id = ?", uploadID).First(&session).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errNotFound("上传会话不存在")
}
if err != nil {
return nil, errInternal("查询上传会话失败: " + err.Error())
}
if session.IsExpired(time.Now()) {
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error
releaseStorage(ctx, d.DB, "presign:"+uploadID)
return nil, errNotFound("上传会话已过期")
}
if expectedMode != "" && session.Mode != expectedMode {
return nil, errBadRequest("此会话不支持" + expectedMode + "模式")
}
return &session, nil
}
// ============ POST /presign/upload/init 初始化预签名上传 ============
// presignInitRequest init 请求体(对齐参考 PresignUploadInitRequest)。
type presignInitRequest struct {
FileName string `json:"file_name" form:"file_name"`
FileSize int64 `json:"file_size" form:"file_size"`
ExpireValue int `json:"expire_value" form:"expire_value"`
ExpireStyle string `json:"expire_style" form:"expire_style"`
Code string `json:"code" form:"code"` // v3.1:自定义提取码(init 时校验,完成时落库)
}
// presignInit 初始化预签名上传(对齐参考 presign_upload_init):
// 引擎支持直链(S3)返回 direct + 预签名 PUT URL;否则返回 proxy + 代理地址。
func (d *Deps) presignInit(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
var req presignInitRequest
if err := bindJSONOrForm(c, &req); err != nil {
respondError(c, err)
return
}
safeName := storage.SanitizeFileName(req.FileName)
if safeName == "" {
auditRecordFailed(c, d.AuditSvc, "文件名非法")
response.Fail(c, http.StatusBadRequest, "文件名非法")
return
}
// v3.1:自定义提取码提前校验(init 时快速失败;完成请求须再次携带)
if err := validatePickupCode(req.Code); err != nil {
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
if err := validateFileMagic(d.Cfg, safeName, "", nil); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件类型被拒绝")
respondError(c, err)
return
}
// v2 需求 ④⑩:动态策略校验(max_file_size0=回落 uploadSize
if err := d.CurrentUploadPolicy().CheckSize(req.FileSize); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小超过限制")
respondError(c, err)
return
}
// M2:直传 confirm 的实际大小校验依赖真实对象,0/负值声明直接拒绝
if req.FileSize <= 0 {
auditRecordFailed(c, d.AuditSvc, "file_size 非法")
response.Fail(c, http.StatusBadRequest, "file_size 必须大于 0")
return
}
if req.ExpireValue <= 0 {
req.ExpireValue = 1
}
if req.ExpireStyle == "" {
req.ExpireStyle = "day"
}
if _, err := resolveExpire(d.Cfg, req.ExpireValue, req.ExpireStyle); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "过期策略非法")
respondError(c, err)
return
}
ctx := c.Request.Context()
uploadID := uuidHex()
resToken := "presign:" + uploadID
if err := reserveStorage(ctx, d.DB, d.Cfg, resToken, req.FileSize, presignSessionExpires*time.Second); err != nil {
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "容量预留失败")
respondError(c, err)
return
}
dirPath, _, _, _, savePath := buildSavePath(d.Cfg, safeName, uploadID)
mode := "proxy"
uploadURL := "/presign/upload/proxy/" + uploadID
putURL, err := d.Store.PresignPutURL(ctx, savePath, presignSessionExpires)
switch {
case err == nil:
mode = "direct"
uploadURL = putURL
case errors.Is(err, storage.ErrNotSupported):
// 引擎不支持直传:代理模式
default:
releaseStorage(ctx, d.DB, resToken)
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "生成预签名失败")
respondError(c, mapStorageError(err))
return
}
session := model.PresignUploadSession{
UploadID: uploadID,
FileName: safeName,
FileSize: req.FileSize,
SavePath: savePath,
Mode: mode,
Engine: d.Store.CurrentName(), // v3:会话归属引擎
ExpireValue: req.ExpireValue,
ExpireStyle: req.ExpireStyle,
ExpiresAt: time.Now().Add(presignSessionExpires * time.Second),
}
if err := d.DB.WithContext(ctx).Create(&session).Error; err != nil {
releaseStorage(ctx, d.DB, resToken)
auditUploadEntry(c, "", safeName, req.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "会话创建失败")
respondError(c, errInternal("创建上传会话失败: "+err.Error()))
return
}
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, uploadID, safeName, req.FileSize, 0)
auditRecordSuccess(c, d.AuditSvc)
detail := gin.H{
"upload_id": uploadID,
"upload_url": uploadURL,
"mode": mode,
"expires_in": presignSessionExpires,
"file_path": dirPath,
}
if mode == "proxy" {
detail["proxy_upload_url"] = uploadURL
detail["legacy_proxy_upload_url"] = "/api" + uploadURL
}
response.OK(c, detail)
}
// ============ PUT /presign/upload/proxy/{uploadID} 代理上传 ============
// presignProxy 代理模式上传(对齐参考 presign_upload_proxy):
// 服务器接收文件并转存到存储引擎,随后立即创建分享记录。
func (d *Deps) presignProxy(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
uploadID := c.Param("uploadID")
session, err := d.getValidSession(c, uploadID, "proxy")
if err != nil {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
// v3.1:自定义提取码随代理上传表单携带(init 时已预校验)
if err := validatePickupCode(c.PostForm("code")); err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
ctx := c.Request.Context()
if err := reserveStorage(ctx, d.DB, d.Cfg, "presign:"+uploadID, session.FileSize, presignSessionExpires*time.Second); err != nil {
respondError(c, err)
return
}
fh, err := c.FormFile("file")
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "缺少 file 字段")
response.Fail(c, http.StatusBadRequest, "缺少上传文件 file 字段")
return
}
// 动态策略快照(与 share/chunk 上传路径一致,消除会话窗口内的策略滞后)
maxSize := d.CurrentUploadPolicy().MaxFileSize
if maxSize > 0 && fh.Size > maxSize {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "大小超过限制")
response.Fail(c, http.StatusForbidden, fmt.Sprintf("大小超过限制,最大为%s", humanSize(maxSize)))
return
}
// 文件大小与声明不符(±1KB 容差,对齐参考)
if abs64(fh.Size-session.FileSize) > 1024 {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小与声明不符")
response.Fail(c, http.StatusBadRequest, "文件大小与声明不符")
return
}
f, err := fh.Open()
if err == nil {
defer func() { _ = f.Close() }()
if err = validateFileMagic(d.Cfg, session.FileName, fh.Header.Get("Content-Type"), readMultipartHeader(f, 64)); err == nil {
// v3:落盘走会话归属引擎
var ps storage.Storage
ps, sErr := d.storeFor(session.Engine)
if sErr != nil {
err = sErr
} else if _, err = ps.SaveFile(ctx, f, session.SavePath); err != nil {
// 落盘失败,err 交给统一错误处理
}
}
}
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件保存失败")
if isStorageErr(err) {
respondError(c, mapStorageError(err))
} else {
respondError(c, errInternal("文件保存失败: "+err.Error()))
}
return
}
code, err := d.createRecordFromSession(c, session, c.PostForm("code"))
releaseStorage(ctx, d.DB, "presign:"+uploadID)
if err != nil {
if ps, sErr := d.storeFor(session.Engine); sErr == nil {
_ = ps.DeleteFile(ctx, session.SavePath) // v3:清理走归属引擎
}
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "创建分享失败")
respondError(c, err)
return
}
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, session.FileName, session.FileSize, fh.Size)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": session.FileName})
}
// ============ POST /presign/upload/confirm/{uploadID} 直传确认 ============
// presignConfirm 直传确认(对齐参考 presign_upload_confirm):
// 客户端完成 S3 直传后调用,校验文件已存在并创建分享记录。
func (d *Deps) presignConfirm(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
uploadID := c.Param("uploadID")
session, err := d.getValidSession(c, uploadID, "direct")
if err != nil {
auditUploadEntry(c, uploadID, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
ctx := c.Request.Context()
if err := reserveStorage(ctx, d.DB, d.Cfg, "presign:"+uploadID, session.FileSize, presignSessionExpires*time.Second); err != nil {
respondError(c, err)
return
}
// v3:直传文件存在性按会话归属引擎检查(直传可能落在旧引擎)
psCheck, sErr := d.storeFor(session.Engine)
if sErr != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "存储引擎不可用: "+sErr.Error())
respondError(c, mapStorageError(sErr))
return
}
// v3.1:自定义提取码随确认请求携带(query 或 JSON/form body,均可选)
customCode := c.Query("code")
if customCode == "" && c.Request.Body != nil && c.Request.ContentLength != 0 {
var fin struct {
Code string `json:"code" form:"code"`
}
if err := bindJSONOrForm(c, &fin); err != nil {
respondError(c, err)
return
}
customCode = fin.Code
}
exists, err := psCheck.FileExists(ctx, session.SavePath)
if err == nil && !exists {
err = errNotFound("文件未上传或上传失败")
}
if err != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件未上传或上传失败")
respondError(c, err)
return
}
// M2 修复:直传内容不经过服务器,confirm 必须核实实际大小与内容类型。
meta, head, hErr := psCheck.HeadMeta(ctx, session.SavePath, 64)
if hErr != nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件信息读取失败")
respondError(c, mapStorageError(hErr))
return
}
if meta == nil {
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件未上传或上传失败")
respondError(c, errNotFound("文件未上传或上传失败"))
return
}
// 大小上限:实际大小超过策略上限 → 删除对象并 403(防绕过 max_file_size / storageLimit
maxSize := d.CurrentUploadPolicy().MaxFileSize
if maxSize > 0 && meta.Size > maxSize {
_ = psCheck.DeleteFile(ctx, session.SavePath)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "实际文件大小超过限制")
respondError(c, errForbidden(fmt.Sprintf("大小超过限制,最大为%s", humanSize(maxSize))))
return
}
// 大小与声明不符(±1KB 容差,对齐 proxy 模式):超差删除对象并 400
if abs64(meta.Size-session.FileSize) > 1024 {
_ = psCheck.DeleteFile(ctx, session.SavePath)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件大小与声明不符")
respondError(c, errBadRequest("文件大小与声明不符"))
return
}
// 内容类型防伪造(对齐 proxy 模式 magic bytes 校验)
if err := validateFileMagic(d.Cfg, session.FileName, "", head); err != nil {
_ = psCheck.DeleteFile(ctx, session.SavePath)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "文件内容校验失败")
respondError(c, err)
return
}
code, err := d.createRecordFromSession(c, session, customCode)
releaseStorage(ctx, d.DB, "presign:"+uploadID)
if err != nil {
if ps, sErr := d.storeFor(session.Engine); sErr == nil {
_ = ps.DeleteFile(ctx, session.SavePath) // v3:清理走归属引擎
}
auditUploadEntry(c, uploadID, session.FileName, session.FileSize, 0)
auditRecordFailed(c, d.AuditSvc, "创建分享失败")
respondError(c, err)
return
}
_ = d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, session.FileName, session.FileSize, session.FileSize)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": session.FileName})
}
// createRecordFromSession 依据预签名会话创建分享记录(对齐参考 create_file_record)。
func (d *Deps) createRecordFromSession(c *gin.Context, session *model.PresignUploadSession, customCode string) (string, error) {
exp, err := resolveExpire(d.Cfg, session.ExpireValue, session.ExpireStyle)
if err != nil {
return "", err
}
// v3.1:完成请求的自定义提取码兜底校验(init 已验,防只发完成请求绕过)
if err := validatePickupCode(customCode); err != nil {
return "", err
}
ctx := c.Request.Context()
code, err := pickCustomCode(ctx, d.DB, d.Cfg, customCode)
if err != nil {
return "", err
}
dir, name := splitDirBase(session.SavePath)
ext := baseExt(name)
fc := model.FileCodes{
Code: code,
Prefix: trimExt(name),
Suffix: ext,
UUIDFileName: &name,
FilePath: &dir,
Size: session.FileSize,
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: session.Engine, // v3:归属引擎戳
}
if err := d.DB.WithContext(ctx).Create(&fc).Error; err != nil {
return "", mapCodeConflict(err) // v3.1:并发占用自定义码 → 友好 400
}
return code, nil
}
// ============ GET /presign/upload/status/{uploadID} ============
// presignStatus 查询预签名会话状态(对齐参考 presign_upload_status)。
func (d *Deps) presignStatus(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
ctx := c.Request.Context()
var session model.PresignUploadSession
if err := d.DB.WithContext(ctx).
Where("upload_id = ?", uploadID).First(&session).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
response.Fail(c, http.StatusNotFound, "上传会话不存在")
return
}
respondError(c, errInternal("查询上传会话失败: "+err.Error()))
return
}
response.OK(c, gin.H{
"upload_id": session.UploadID,
"file_name": session.FileName,
"file_size": session.FileSize,
"mode": session.Mode,
"created_at": session.CreatedAt.Format(time.RFC3339),
"expires_at": session.ExpiresAt.Format(time.RFC3339),
"is_expired": session.IsExpired(time.Now()),
})
}
// ============ DELETE /presign/upload/{uploadID} 取消会话 ============
// presignCancel 取消预签名上传会话(对齐参考 presign_upload_cancel):
// 直传模式尽力清理已直传的文件。
func (d *Deps) presignCancel(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
uploadID := c.Param("uploadID")
ctx := c.Request.Context()
session, err := d.getValidSession(c, uploadID, "")
if err != nil {
respondError(c, err)
return
}
if session.Mode == "direct" {
// v3:清理走会话归属引擎
if ps, sErr := d.storeFor(session.Engine); sErr == nil {
if exists, eErr := ps.FileExists(ctx, session.SavePath); eErr == nil && exists {
_ = ps.DeleteFile(ctx, session.SavePath)
}
}
}
if err := d.DB.WithContext(ctx).Where("upload_id = ?", uploadID).
Delete(&model.PresignUploadSession{}).Error; err != nil {
respondError(c, errInternal("取消上传会话失败: "+err.Error()))
return
}
releaseStorage(ctx, d.DB, "presign:"+uploadID)
response.OK(c, gin.H{"message": "上传会话已取消"})
}
// ============ 杂项 ============
// abs64 绝对值。
func abs64(n int64) int64 {
if n < 0 {
return -n
}
return n
}
// isStorageErr 判断是否为存储层哨兵错误(含 %w 包装)。
func isStorageErr(err error) bool {
return err != nil && (errors.Is(err, storage.ErrNotFound) ||
errors.Is(err, storage.ErrInvalidPath) ||
errors.Is(err, storage.ErrUnavailable) ||
errors.Is(err, storage.ErrNotSupported) ||
errors.Is(err, storage.ErrRangeNotSatisfiable) ||
errors.Is(err, storage.ErrHashMismatch))
}
+187
View File
@@ -0,0 +1,187 @@
package api
import (
"net/http"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/audit"
"filecodebox/internal/config"
"filecodebox/internal/middleware"
"filecodebox/internal/response"
"filecodebox/internal/settings"
"filecodebox/internal/storage"
)
// Deps API 层共享依赖(main.go 装配后注入)。
type Deps struct {
DB *gorm.DB
Cfg *config.Config
Mgr *settings.Manager
AuditSvc *audit.Service
Limiter *middleware.RateLimiter
Store *storage.Manager // v3:可热切换引擎管理器(实现 Storage 接口)
Version string
}
// jwtSecret 当前 JWT 签名密钥(settings KV 运行时可变)。
func (d *Deps) jwtSecret() string { return d.Mgr.SecretProvider()() }
// Register 注册全部 API 路由与前端静态资源回退。
// 业务路由挂根路径(/share /chunk /presign /admin),与审计中间件
// DefaultClassifier 的路由模式一致(t1 冻结契约);公共接口保留
// /api/v1/health 与 /api/v1/config(对齐 t1 骨架)。
func Register(r *gin.Engine, d *Deps) {
// —— 公共接口 ——
r.GET("/api/v1/health", d.health)
r.GET("/api/v1/config", d.publicConfig)
// Info3robotsText 配置键此前无路由承接,补上(内容可在管理端自定义)
r.GET("/robots.txt", d.robotsText)
// —— 初始化向导(未初始化时唯一可用入口,GuardNotInitialized 白名单)——
registerSetup(r, d)
// —— 分享 ——
share := r.Group("/share")
{
share.POST("/text", d.shareText)
share.POST("/file", d.shareFile)
// metadata:每次访问即计数(RequireRateLimit=进入检查+完成计数)
share.GET("/metadata", d.Limiter.RequireRateLimit(middleware.LimitMeta), d.shareMetadata)
share.POST("/metadata", d.Limiter.RequireRateLimit(middleware.LimitMeta), d.shareMetadataPost)
share.GET("/select", d.shareSelect)
share.POST("/select", d.shareSelectPost)
share.GET("/download", d.shareDownload)
}
// —— 分片上传 ——
chunk := r.Group("/chunk")
{
chunk.POST("/upload/init", d.chunkInit)
// 主路径(参考语义):/chunk/upload/{uploadID}/{index}
// 扁平兼容:/chunk/upload + 表单/query 传 upload_id/chunk_index
chunk.POST("/upload/:uploadID/:chunkIndex", d.chunkUpload)
chunk.POST("/upload", d.chunkUploadFlat)
chunk.GET("/upload/status/:uploadID", d.chunkStatus)
chunk.POST("/upload/complete/:uploadID", d.chunkComplete)
chunk.DELETE("/upload/:uploadID", d.chunkCancel)
}
// —— 预签名直传 ——
presign := r.Group("/presign")
{
presign.POST("/upload/init", d.presignInit)
presign.PUT("/upload/proxy/:uploadID", d.presignProxy)
presign.POST("/upload/confirm/:uploadID", d.presignConfirm)
presign.GET("/upload/status/:uploadID", d.presignStatus)
presign.DELETE("/upload/:uploadID", d.presignCancel)
}
// —— 管理端(login 公开,其余需管理员 JWT)——
registerAdmin(r, d)
// —— 前端静态资源 + SPA 回退(须最后注册)——
registerWeb(r, d)
}
// health 健康检查(对齐 t1 骨架,保持 /api/v1/health 语义不变)。
func (d *Deps) health(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"code": 200,
"msg": "ok",
"data": gin.H{
"status": "ok",
"version": d.Version,
"storage": d.Cfg.Engine(),
"time": nowRFC3339(),
},
})
}
// robotsText 输出管理端可配置的 robots.txt 内容。
func (d *Deps) robotsText(c *gin.Context) {
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(d.Cfg.GetString("robotsText")))
}
// publicConfig 公共配置(前端首页/上传页所需;v2 需求 ①②③④⑩ 扩展):
// - 展示字段:站点信息、Logo/favicon、背景图、页脚文案/备案号、通知;
// - 策略范围(上传页动态渲染):大小上限、类型白名单、过期方式、保存
// 时间/次数上限、上传频率(仅范围,不含内部实现键)。
//
// 敏感键(admin_token/jwt_secretsettings.SensitiveKeys)与本端点无关:
// 下发字段为白名单显式构造,任何敏感键均不会出现在响应中。
func (d *Deps) publicConfig(c *gin.Context) {
cfg := d.Cfg
policy := d.CurrentUploadPolicy()
uploadCount := cfg.GetInt("uploadCount")
uploadMinute := cfg.GetInt("uploadMinute")
// uploadSize 为参考语义的回落上限,单独下发供管理端联动展示
c.JSON(http.StatusOK, gin.H{
"code": 200,
"msg": "ok",
"data": gin.H{
"config": gin.H{
"name": cfg.SiteName(),
"description": cfg.GetString("description"),
"explain": cfg.GetString("page_explain"),
// 需求 ①:Logo/favicon/背景图
"logo_url": cfg.LogoURL(),
"favicon_url": cfg.FaviconURL(),
"background_url": cfg.BackgroundURL(),
// 需求 ②:页脚自定义内容与备案号
"footer_text": cfg.FooterText(),
"footer_beian": cfg.FooterBeian(),
// v3:当前存储引擎名(仅名称,任何引擎参数/凭据不下发)
"storage_engine": d.Store.CurrentName(),
"site_domain": d.Cfg.SiteDomain(),
// 需求 ③:系统通知(开关 + 内容,前台右上角悬浮窗)
// L7:读取侧再做一次白名单净化,覆盖历史存量与直改库的数据
"notify_enabled": boolToInt(cfg.NotifyEnabled()),
"notify_title": cfg.GetString("notify_title"),
"notify_content": settings.SanitizeInlineHTML(cfg.GetString("notify_content")),
// 策略范围(需求 ④⑩):上传页动态读取并在范围内选择
"uploadSize": cfg.UploadSize(),
"max_file_size": policy.MaxFileSize,
"maxFileSize": policy.MaxFileSize,
"allowedFileTypes": policy.AllowedTypes,
"expireStyle": policy.ExpireStyles,
"max_save_seconds": policy.MaxSaveSeconds,
"maxSaveSeconds": policy.MaxSaveSeconds,
"max_save_count": policy.MaxSaveCount,
"maxSaveCount": policy.MaxSaveCount,
"uploadCount": uploadCount,
"uploadMinute": uploadMinute,
"enableChunk": cfg.EnableChunk(),
"openUpload": cfg.OpenUpload(),
},
"meta": gin.H{
"version": d.Version,
"features": gin.H{
"chunkUpload": cfg.EnableChunk(),
"guestUpload": cfg.OpenUpload(),
},
},
},
})
}
// requireShareLogin 分享上传权限(对齐参考 share_required_login):
// openUpload 开启时游客可传;关闭时要求管理员 Bearer token403)。
func (d *Deps) requireShareLogin(c *gin.Context) bool {
if d.Cfg.OpenUpload() {
return true
}
header := c.GetHeader("Authorization")
const prefix = "Bearer "
if len(header) <= len(prefix) || header[:len(prefix)] != prefix {
response.Fail(c, http.StatusForbidden, "本站未开启游客上传,如需上传请先登录后台")
return false
}
token := header[len(prefix):]
if _, err := middleware.VerifyAdminToken(d.jwtSecret(), token); err != nil {
response.Fail(c, http.StatusForbidden, "本站未开启游客上传,如需上传请先登录后台")
return false
}
return true
}
+193
View File
@@ -0,0 +1,193 @@
// security_fixes_test.go — 安全审计修复项行为测试:
// L4 enableChunk 强制、M2 presign 大小/类型校验、L3 提码长度、M3 chunk_size 上限。
package api
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/model"
"filecodebox/internal/settings"
)
// postJSON 以 JSON body 调用 POST 端点。
func postJSON(d *Deps, path string, body any) *httptest.ResponseRecorder {
var reader *bytes.Reader
if body == nil {
reader = bytes.NewReader(nil)
} else {
raw, _ := json.Marshal(body)
reader = bytes.NewReader(raw)
}
req := httptest.NewRequest(http.MethodPost, path, reader)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = req
c.Params = append(c.Params, gin.Param{Key: "uploadID", Value: req.URL.Path[len("/presign/upload/confirm/"):]})
d.presignConfirm(c)
return w
}
// sha256LegacyHash 构造旧版 sha256$salt$hash 格式(M1 迁移测试用)。
func sha256LegacyHash(password string) string {
salt := make([]byte, 16)
for i := range salt {
salt[i] = byte(i)
}
saltHex := hex.EncodeToString(salt)
sum := sha256.Sum256([]byte(saltHex + password))
return "sha256$" + saltHex + "$" + hex.EncodeToString(sum[:])
}
// TestChunkToggleEnforced L4enableChunk=0 时 /chunk 相关端点一律 403。
func TestChunkToggleEnforced(t *testing.T) {
d := newPolicyTestDeps(t)
// 默认 enableChunk=0
w := chunkInitJSON(d, `{"file_name":"a.png","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusForbidden {
t.Fatalf("enableChunk=0 时 init 应 403: %d %s", code, w.Body.String())
}
// 开启后放行
if w := patchConfig(d, map[string]any{"enableChunk": 1}); w.Code != 200 {
t.Fatalf("patch enableChunk: %d", w.Code)
}
w = chunkInitJSON(d, `{"file_name":"a.png","file_size":500,"chunk_size":1024,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != 200 {
t.Fatalf("enableChunk=1 时 init 应 200: %d %s", code, w.Body.String())
}
}
// TestChunkSizeCap M3chunk_size 超过 32MB 上限时 400。
func TestChunkSizeCap(t *testing.T) {
d := newPolicyTestDeps(t)
if w := patchConfig(d, map[string]any{"enableChunk": 1}); w.Code != 200 {
t.Fatalf("patch enableChunk: %d", w.Code)
}
w := chunkInitJSON(d, `{"file_name":"a.bin","file_size":70000000000,"chunk_size":34000000,"file_hash":"h"}`)
if code, _ := respBody(t, w); code != http.StatusBadRequest {
t.Fatalf("chunk_size 超上限应 400: %d %s", code, w.Body.String())
}
}
// TestPickupCodeMinLen L34 位自定义码拒绝、5 位通过。
func TestPickupCodeMinLen(t *testing.T) {
if err := validatePickupCode("abcd"); err == nil {
t.Fatal("4 位码应被拒绝")
}
if err := validatePickupCode("abcde"); err != nil {
t.Fatalf("5 位码应通过: %v", err)
}
}
// TestPresignConfirmRejectsOversizeObject M2
// 直传会话 confirm 时,若对象实际大小超过策略上限,应删除对象并 403。
func TestPresignConfirmRejectsOversizeObject(t *testing.T) {
d := newPolicyTestDeps(t)
ctx := context.Background()
// 声明 10 字节、策略上限 100 → 实际 PUT 500 字节对象
if err := d.Mgr.UpdateKV(ctx, map[string]any{"max_file_size": 100}); err != nil {
t.Fatalf("UpdateKV: %v", err)
}
if err := d.Mgr.Reload(ctx); err != nil {
t.Fatalf("Reload: %v", err)
}
uploadID := "test-oversize-confirm"
savePath := "share/data/presign_test.bin"
if _, err := d.Store.SaveFile(ctx, bytes.NewReader(make([]byte, 500)), savePath); err != nil {
t.Fatalf("SaveFile: %v", err)
}
sess := model.PresignUploadSession{
UploadID: uploadID, FileName: "presign_test.bin", FileSize: 10,
SavePath: savePath, Mode: "direct",
ExpireValue: 1, ExpireStyle: "day",
CreatedAt: time.Now(), ExpiresAt: time.Now().Add(time.Hour), Engine: "local",
}
if err := d.DB.WithContext(ctx).Create(&sess).Error; err != nil {
t.Fatalf("create session: %v", err)
}
res := model.StorageReservation{Token: "presign:" + uploadID, Size: 10, ExpiresAt: time.Now().Add(time.Hour)}
if err := d.DB.WithContext(ctx).Create(&res).Error; err != nil {
t.Fatalf("create reservation: %v", err)
}
w := postJSON(d, "/presign/upload/confirm/"+uploadID, nil)
if w.Code != http.StatusForbidden {
t.Fatalf("超限对象 confirm 应 403: %d %s", w.Code, w.Body.String())
}
// 对象应被删除、预留应释放
if ok, _ := d.Store.FileExists(ctx, savePath); ok {
t.Fatal("超限对象应被服务端删除")
}
var cnt int64
_ = d.DB.WithContext(ctx).Model(&model.StorageReservation{}).Where("token = ?", res.Token).Count(&cnt).Error
if cnt != 0 {
t.Fatal("预留应被释放")
}
}
// TestPresignConfirmRejectsSizeMismatch M2:实际大小与声明差超过 ±1KB 时 400。
func TestPresignConfirmRejectsSizeMismatch(t *testing.T) {
d := newPolicyTestDeps(t)
ctx := context.Background()
uploadID := "test-mismatch-confirm"
savePath := "share/data/presign_mismatch.bin"
if _, err := d.Store.SaveFile(ctx, bytes.NewReader(make([]byte, 2048)), savePath); err != nil {
t.Fatalf("SaveFile: %v", err)
}
sess := model.PresignUploadSession{
UploadID: uploadID, FileName: "presign_mismatch.bin", FileSize: 10,
SavePath: savePath, Mode: "proxy", // proxy 模式同样走大小核对(多引擎一致)
ExpireValue: 1, ExpireStyle: "day",
CreatedAt: time.Now(), ExpiresAt: time.Now().Add(time.Hour), Engine: "local",
}
if err := d.DB.WithContext(ctx).Create(&sess).Error; err != nil {
t.Fatalf("create session: %v", err)
}
res := model.StorageReservation{Token: "presign:" + uploadID, Size: 10, ExpiresAt: time.Now().Add(time.Hour)}
if err := d.DB.WithContext(ctx).Create(&res).Error; err != nil {
t.Fatalf("create reservation: %v", err)
}
w := postJSON(d, "/presign/upload/confirm/"+uploadID, nil)
if w.Code != http.StatusBadRequest {
t.Fatalf("大小不符 confirm 应 400: %d %s", w.Code, w.Body.String())
}
}
// TestAdminPasswordAutoUpgrade M1:明文/旧哈希经 VerifyPassword 后 NeedsRehash 为真,
// bcrypt 哈希不再需要升级。
func TestAdminPasswordAutoUpgrade(t *testing.T) {
if !settings.NeedsRehash("FileCodeBox2023") {
t.Fatal("明文哈希需要升级")
}
legacy := sha256LegacyHash("pwd12345")
if !settings.NeedsRehash(legacy) {
t.Fatal("sha256 哈希需要升级")
}
if !settings.VerifyPassword("pwd12345", legacy) {
t.Fatal("旧 sha256 哈希兼容校验失败")
}
b := settings.HashPassword("pwd12345")
if settings.NeedsRehash(b) {
t.Fatal("bcrypt 哈希不需要升级")
}
if !settings.VerifyPassword("pwd12345", b) {
t.Fatal("bcrypt 校验失败")
}
if settings.VerifyPassword("wrong", b) {
t.Fatal("错误密码不应通过")
}
}
+327
View File
@@ -0,0 +1,327 @@
package api
import (
"net/http"
"strconv"
"strings"
"github.com/gin-gonic/gin"
"filecodebox/internal/response"
"filecodebox/internal/settings"
)
// fileSizeUnits 文件大小单位(对齐参考 FILE_SIZE_UNITS)。
var fileSizeUnits = map[string]int64{"KB": 1024, "MB": 1024 * 1024, "GB": 1024 * 1024 * 1024}
// saveTimeUnits 保存时间单位(秒)。
var saveTimeUnits = map[string]int64{"second": 1, "minute": 60, "hour": 3600, "day": 86400}
// expireStyleOptions 可用过期方式(用于 setup 表单校验)。
var expireStyleOptions = []string{"day", "hour", "minute", "forever", "count"}
// setupFormValue 取表单/JSON 字符串值。
func setupFormValue(data map[string]any, key, def string) string {
v, ok := data[key]
if !ok || v == nil {
return def
}
if s, ok := v.(string); ok {
return s
}
return strings.TrimSpace(strconv.FormatFloat(toAnyFloat(v), 'f', -1, 64))
}
func toAnyFloat(v any) float64 {
switch n := v.(type) {
case float64:
return n
case int:
return float64(n)
case int64:
return float64(n)
}
return 0
}
// registerSetup 注册初始化向导(未初始化时唯一可用入口,白名单 /setup)。
func registerSetup(r *gin.Engine, d *Deps) {
r.GET("/setup", func(c *gin.Context) {
if d.Mgr.IsInitialized() {
c.Redirect(http.StatusSeeOther, "/")
return
}
c.Data(http.StatusOK, "text/html; charset=utf-8", []byte(buildSetupPage("")))
})
r.POST("/setup", func(c *gin.Context) {
if d.Mgr.IsInitialized() {
c.Redirect(http.StatusSeeOther, "/")
return
}
d.setupSubmit(c)
})
}
// setupSubmit 处理初始化提交(对齐参考 setup_submit + parse_setup_options)。
func (d *Deps) setupSubmit(c *gin.Context) {
// 兼容 JSON 与表单
data := map[string]any{}
if strings.Contains(c.GetHeader("Content-Type"), "application/json") {
if err := c.ShouldBindJSON(&data); err != nil {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("请求体格式错误")))
return
}
} else if err := c.Request.ParseForm(); err != nil {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("请求体格式错误")))
return
} else {
// 多值字段(如多个 expireStyle 复选框)保留完整列表,单值取首项
for k, v := range c.Request.PostForm {
switch {
case len(v) == 1:
data[k] = v[0]
case len(v) > 1:
data[k] = v
}
}
}
adminPassword := setupFormValue(data, "admin_password", "")
confirmPassword := setupFormValue(data, "confirm_password", "")
siteName := setupFormValue(data, "site_name", "")
if adminPassword == "" || len(adminPassword) < 8 {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("管理员密码至少 8 位")))
return
}
if adminPassword != confirmPassword {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage("两次输入的管理员密码不一致")))
return
}
patch, errMsg := parseSetupOptions(data)
if errMsg != "" {
c.Data(http.StatusBadRequest, "text/html; charset=utf-8", []byte(buildSetupPage(errMsg)))
return
}
patch["site_name"] = firstNonEmpty(siteName, "文件快传")
patch["admin_token"] = settings.HashPassword(adminPassword)
patch["jwt_secret"] = settings.GenerateJWTSecret()
ctx := c.Request.Context()
if err := d.Mgr.UpdateKV(ctx, patch); err != nil {
c.Data(http.StatusInternalServerError, "text/html; charset=utf-8", []byte(buildSetupPage("初始化失败: "+err.Error())))
return
}
if err := d.Mgr.Reload(ctx); err != nil {
c.Data(http.StatusInternalServerError, "text/html; charset=utf-8", []byte(buildSetupPage("配置重载失败: "+err.Error())))
return
}
d.syncRateRules()
// JSON 请求返回 JSON;表单返回成功页
if strings.Contains(c.GetHeader("Accept"), "application/json") ||
strings.Contains(c.GetHeader("Content-Type"), "application/json") {
response.OK(c, gin.H{"ok": true, "admin": "/#/admin"})
return
}
c.Data(http.StatusOK, "text/html; charset=utf-8", []byte(buildSetupSuccessPage()))
}
// parseSetupOptions 解析并校验初始化选项(对齐参考 parse_setup_options)。
func parseSetupOptions(data map[string]any) (map[string]any, string) {
out := map[string]any{}
// 文件大小限制
unit := strings.ToUpper(setupFormValue(data, "upload_size_unit", "MB"))
if _, ok := fileSizeUnits[unit]; !ok {
return nil, "文件大小单位不正确"
}
sizeVal, err := strconv.Atoi(setupFormValue(data, "upload_size_value", "10"))
if err != nil || sizeVal < 1 {
return nil, "文件大小限制必须是正整数"
}
out["uploadSize"] = int64(sizeVal) * fileSizeUnits[unit]
// 最长保存时间
saveUnit := strings.ToLower(setupFormValue(data, "save_time_unit", "day"))
if _, ok := saveTimeUnits[saveUnit]; !ok {
return nil, "最长保存时间单位不正确"
}
saveVal, err := strconv.Atoi(setupFormValue(data, "save_time_value", "0"))
if err != nil || saveVal < 0 {
return nil, "最长保存时间必须是非负整数"
}
out["max_save_seconds"] = int64(saveVal) * saveTimeUnits[saveUnit]
// 过期方式白名单
var styles []string
if raw, ok := data["expireStyle"]; ok {
switch v := raw.(type) {
case []any:
for _, item := range v {
if s, ok := item.(string); ok {
styles = append(styles, s)
}
}
case []string:
styles = v
case string:
for _, s := range strings.Split(v, ",") {
styles = append(styles, strings.TrimSpace(s))
}
}
}
valid := map[string]bool{}
var finalStyles []string
for _, s := range styles {
s = strings.TrimSpace(s)
if s == "" || valid[s] {
continue
}
for _, opt := range expireStyleOptions {
if opt == s {
valid[s] = true
finalStyles = append(finalStyles, s)
break
}
}
}
if len(finalStyles) == 0 {
return nil, "至少需要选择一种过期方式"
}
out["expireStyle"] = finalStyles
// 取件码类型
codeType := setupFormValue(data, "code_generate_type", "secret")
if codeType != "number" && codeType != "secret" {
return nil, "提取码类型不正确"
}
out["code_generate_type"] = codeType
// 频率限制
for _, item := range []struct{ key, def string }{
{"errorCount", "10"}, {"errorMinute", "1"},
{"loginCount", "5"}, {"loginMinute", "15"},
{"uploadCount", "10"}, {"uploadMinute", "1"},
} {
n, err := strconv.Atoi(setupFormValue(data, item.key, item.def))
if err != nil || n < 1 {
return nil, item.key + " 必须是正整数"
}
out[item.key] = n
}
// 布尔开关
out["openUpload"] = boolToInt(parseSetupBool(data, "openUpload", true))
out["enableChunk"] = boolToInt(parseSetupBool(data, "enableChunk", false))
// 允许文件类型
allowed := setupFormValue(data, "allowed_file_types", "*")
var types []string
for _, item := range strings.Split(allowed, ",") {
if item = strings.TrimSpace(item); item != "" {
types = append(types, item)
}
}
if len(types) == 0 {
types = []string{"*"}
}
out["allowed_file_types"] = types
return out, ""
}
// parseSetupBool 解析表单布尔(缺省 default"1"/"true"/"on"/"yes" 为真)。
func parseSetupBool(data map[string]any, key string, def bool) bool {
v, ok := data[key]
if !ok {
return def
}
switch s := v.(type) {
case string:
switch strings.ToLower(strings.TrimSpace(s)) {
case "1", "true", "on", "yes":
return true
case "0", "false", "off", "no", "":
return false
}
case float64:
return s != 0
case bool:
return s
}
return def
}
// buildSetupPage 初始化向导页面(简洁中文表单)。
func buildSetupPage(errMsg string) string {
errBlock := ""
if errMsg != "" {
errBlock = `<div style="margin-bottom:12px;padding:10px 12px;border-radius:10px;background:#fef2f2;color:#b91c1c;font-size:13px">` + htmlEscape(errMsg) + `</div>`
}
return `<!doctype html>
<html lang="zh-CN">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>初始化 文件快传</title>
<style>
body{margin:0;min-height:100vh;display:grid;place-items:center;padding:16px;font-family:-apple-system,"Segoe UI",sans-serif;background:#f5f5f7;color:#18181b}
main{width:min(100%,640px);padding:24px;border-radius:16px;background:#fff;box-shadow:0 18px 50px rgba(23,32,51,.08)}
h1{margin:0 0 6px;font-size:20px} p{margin:0 0 16px;color:#71717a;font-size:13px}
label{display:block;margin:10px 0 4px;font-size:12px;color:#3f3f46;font-weight:600}
input{width:100%;height:36px;border:1px solid #e4e4e7;border-radius:8px;padding:0 10px;box-sizing:border-box;font:inherit}
.grid{display:grid;grid-template-columns:1fr 1fr;gap:0 12px}
button{width:100%;height:40px;margin-top:16px;border:0;border-radius:10px;background:#18181b;color:#fff;font:inherit;font-weight:700;cursor:pointer}
</style>
</head>
<body><main>
<h1>初始化 文件快传</h1>
<p>首次配置管理员密码、上传限制和取件策略,后续可在后台调整。</p>
` + errBlock + `
<form method="post" action="/setup" autocomplete="off">
<label>站点名称</label>
<input name="site_name" maxlength="80" placeholder="文件快传">
<div class="grid">
<div><label>管理员密码</label><input name="admin_password" type="password" minlength="8" required></div>
<div><label>确认管理员密码</label><input name="confirm_password" type="password" minlength="8" required></div>
<div><label>单文件大小限制</label><input name="upload_size_value" type="number" min="1" value="10" required></div>
<div><label>大小单位</label><input name="upload_size_unit" value="MB" required></div>
<div><label>上传频率(次/分钟)</label><input name="uploadCount" type="number" min="1" value="10" required></div>
<div><label>上传检测窗口(分钟)</label><input name="uploadMinute" type="number" min="1" value="1" required></div>
<div><label>取件错误频率(次/分钟)</label><input name="errorCount" type="number" min="1" value="10" required></div>
<div><label>取件错误窗口(分钟)</label><input name="errorMinute" type="number" min="1" value="1" required></div>
<div><label>登录失败频率(次/分钟)</label><input name="loginCount" type="number" min="1" value="5" required></div>
<div><label>登录失败窗口(分钟)</label><input name="loginMinute" type="number" min="1" value="15" required></div>
<div><label>最长保存时间</label><input name="save_time_value" type="number" min="0" value="0" required></div>
<div><label>保存时间单位</label><input name="save_time_unit" value="day" required></div>
</div>
<label>允许文件类型(逗号分隔,* 不限制)</label>
<input name="allowed_file_types" value="*">
<label>提取码类型(number=数字 / secret=随机字符)</label>
<input name="code_generate_type" value="secret">
<label><input type="checkbox" name="openUpload" value="1" checked style="width:auto"> 允许游客上传</label>
<label><input type="checkbox" name="enableChunk" value="1" style="width:auto"> 启用切片上传</label>
<input type="hidden" name="expireStyle" value="day">
<input type="hidden" name="expireStyle" value="hour">
<input type="hidden" name="expireStyle" value="minute">
<input type="hidden" name="expireStyle" value="forever">
<input type="hidden" name="expireStyle" value="count">
<button type="submit">完成初始化</button>
</form>
</main></body></html>`
}
// buildSetupSuccessPage 初始化完成页。
func buildSetupSuccessPage() string {
return `<!doctype html>
<html lang="zh-CN"><head><meta charset="utf-8"><meta http-equiv="refresh" content="2;url=/#/admin"><title>初始化完成</title></head>
<body style="display:grid;place-items:center;min-height:100vh;font-family:-apple-system,sans-serif;background:#f6f8fb;color:#172033">
<main style="text-align:center;padding:32px;background:#fff;border-radius:12px;box-shadow:0 18px 50px rgba(23,32,51,.08)">
<h1>初始化完成</h1><p>管理员密码已设置,请使用刚才的密码登录后台。</p><a href="/#/admin">进入后台</a>
</main></body></html>`
}
// htmlEscape HTML 转义(错误信息拼接用)。
func htmlEscape(s string) string {
r := strings.NewReplacer("&", "&amp;", "<", "&lt;", ">", "&gt;", `"`, "&#39;")
return r.Replace(s)
}
+626
View File
@@ -0,0 +1,626 @@
package api
import (
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
"filecodebox/internal/audit"
"filecodebox/internal/middleware"
"filecodebox/internal/model"
"filecodebox/internal/response"
"filecodebox/internal/storage"
)
// nowRFC3339 当前时间的 RFC3339 表示。
func nowRFC3339() string { return time.Now().Format(time.RFC3339) }
// lookupByCode 按取件码查询分享记录(对齐参考 get_code_file_by_code):
// 不存在返回 "文件不存在"expired=true 时过期返回 "文件已过期"。
// 服务端兜底:历史「链接+提取码」复制格式会把「CODE CODE」整串当码传入
// (空格经 URL 编码进 query/path),取第一段有效码避免误报不存在。
func (d *Deps) lookupByCode(c *gin.Context, code string, checkExpired bool) (*model.FileCodes, error) {
code = strings.TrimSpace(code)
if fields := strings.Fields(code); len(fields) > 1 {
code = fields[0]
}
if code == "" {
return nil, errNotFound("文件不存在")
}
var fc model.FileCodes
err := d.DB.WithContext(c.Request.Context()).
Where("code = ?", code).First(&fc).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errNotFound("文件不存在")
}
if err != nil {
return nil, errInternal("查询失败: " + err.Error())
}
if checkExpired && fc.Expired(time.Now()) {
return nil, errNotFound("文件已过期")
}
return &fc, nil
}
// consumeUsage 原子校验分享状态并记录一次实际领取(对齐参考 consume_file_usage):
// 仅当 expired_count>0(次数剩余)或 expired_count<0 且未到过期时间时扣减成功。
func (d *Deps) consumeUsage(c *gin.Context, fc *model.FileCodes) bool {
now := time.Now()
res := d.DB.WithContext(c.Request.Context()).
Model(&model.FileCodes{}).
Where("id = ?", fc.ID).
Where("expired_count > 0 OR (expired_count < 0 AND (expired_at IS NULL OR expired_at > ?))", now).
Updates(map[string]any{
"expired_count": gorm.Expr("CASE WHEN expired_count > 0 THEN expired_count - 1 ELSE expired_count END"),
"used_count": gorm.Expr("used_count + 1"),
})
return res.Error == nil && res.RowsAffected > 0
}
// fileSavePath 拼接分享记录的存储相对路径(file_path/uuid_file_name)。
func fileSavePath(fc *model.FileCodes) string {
dir := ""
if fc.FilePath != nil {
dir = strings.Trim(*fc.FilePath, "/")
}
name := ""
if fc.UUIDFileName != nil {
name = *fc.UUIDFileName
}
if dir == "" {
return name
}
return dir + "/" + name
}
// ============ POST /share/text 文本分享 ============
// shareText 创建文本分享(对齐参考 share_text)。
func (d *Deps) shareText(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
// v3.1 修复:JSON/表单/ultipart 统一绑定(form+json 双标签——此前仅 PostForm 时,
// JSON 提交会静默存成空文本并 200,取件页空白)。
var body struct {
Text string `json:"text" form:"text"`
ExpireValue int `json:"expire_value" form:"expire_value"`
ExpireStyle string `json:"expire_style" form:"expire_style"`
Code string `json:"code" form:"code"`
}
if err := bindJSONOrForm(c, &body); err != nil {
respondError(c, err)
return
}
text := body.Text
if strings.TrimSpace(text) == "" {
response.Fail(c, http.StatusBadRequest, "分享内容不能为空")
return
}
// M3:前置拒绝超大 body(配合全局 BodyLimit441KB 为 222KB 内容 + 表单/JSON 编码余量)
if c.Request.ContentLength > 441*1024 {
response.Fail(c, http.StatusForbidden, "内容过多,建议采用文件形式")
return
}
expireValue := body.ExpireValue
if expireValue == 0 {
expireValue = 1
}
expireStyle := body.ExpireStyle
if expireStyle == "" {
expireStyle = "day"
}
// v3.1:自定义提取码格式校验(4-8 位字母数字,空=随机)
if err := validatePickupCode(body.Code); err != nil {
respondError(c, err)
return
}
exp, err := resolveExpire(d.Cfg, expireValue, expireStyle)
if err != nil {
respondError(c, err)
return
}
textSize := int64(len([]byte(text)))
if textSize > 222*1024 {
response.Fail(c, http.StatusForbidden, "内容过多,建议采用文件形式")
return
}
ctx := c.Request.Context()
token := "text:" + uuidHex()
if err := reserveStorage(ctx, d.DB, d.Cfg, token, textSize, 300*time.Second); err != nil {
respondError(c, err)
return
}
code, err := pickCustomCode(ctx, d.DB, d.Cfg, body.Code)
if err == nil {
fc := model.FileCodes{
Code: code,
Text: &text,
Size: textSize,
Prefix: "Text",
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: d.Store.CurrentName(), // v3:归属引擎戳(文本也记录,保持一致性)
}
err = d.DB.WithContext(ctx).Create(&fc).Error
}
err = mapCodeConflict(err) // v3.1:自定义码唯一索引冲突 → 友好 400
releaseStorage(ctx, d.DB, token)
if err != nil {
auditRecordFailed(c, d.AuditSvc, "文本分享创建失败")
respondError(c, err)
return
}
d.Limiter.Add(c, middleware.LimitUpload) // 上传成功才计数
auditUploadEntry(c, code, "Text", textSize, textSize)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code})
}
// ============ POST /share/file 文件分享 ============
// shareFile 上传文件并创建分享(对齐参考 share_file)。
func (d *Deps) shareFile(c *gin.Context) {
if !d.requireShareLogin(c) {
return
}
if !requireUploadLimit(c, d.Limiter) {
return
}
fh, err := c.FormFile("file")
if err != nil {
auditRecordFailed(c, d.AuditSvc, "缺少 file 字段")
response.Fail(c, http.StatusBadRequest, "缺少上传文件 file 字段")
return
}
origName := fh.Filename
// v2 需求 ④⑩:动态策略校验(max_file_size0=回落 uploadSize
if err := d.CurrentUploadPolicy().CheckSize(fh.Size); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "大小超过限制")
respondError(c, err)
return
}
expireValue := formInt(c, "expire_value", 1)
expireStyle := c.DefaultPostForm("expire_style", "day")
exp, err := resolveExpire(d.Cfg, expireValue, expireStyle)
if err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "过期策略非法")
respondError(c, err)
return
}
// v3.1:自定义提取码(落盘前校验,失败快速返回不占容量预留)
if err := validatePickupCode(c.PostForm("code")); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "提取码非法")
respondError(c, err)
return
}
// magic bytes 防伪造(读前 64 字节)
f, err := fh.Open()
if err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件读取失败")
response.Fail(c, http.StatusBadRequest, "文件读取失败")
return
}
defer func() { _ = f.Close() }()
if err := validateFileMagic(d.Cfg, origName, fh.Header.Get("Content-Type"), readMultipartHeader(f, 64)); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件类型被拒绝")
respondError(c, err)
return
}
dirPath, prefix, suffix, cleanName, savePath := buildSavePath(d.Cfg, origName, uuidHex())
ctx := c.Request.Context()
resToken := "file:" + uuidHex()
if err := reserveStorage(ctx, d.DB, d.Cfg, resToken, fh.Size, time.Hour); err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "容量预留失败")
respondError(c, err)
return
}
code, err := pickCustomCode(ctx, d.DB, d.Cfg, c.PostForm("code"))
if err == nil {
if _, err = d.Store.SaveFile(ctx, f, savePath); err != nil {
// 保存失败:清理半写文件
_ = d.Store.DeleteFile(ctx, savePath)
}
}
if err == nil {
fc := model.FileCodes{
Code: code,
Prefix: prefix,
Suffix: suffix,
UUIDFileName: &cleanName,
FilePath: &dirPath,
Size: fh.Size,
ExpiredAt: exp.ExpiredAt,
ExpiredCount: exp.ExpiredCount,
UsedCount: exp.UsedCount,
Engine: d.Store.CurrentName(), // v3:归属引擎戳(下载按此取回)
}
if err = d.DB.WithContext(ctx).Create(&fc).Error; err != nil {
err = mapCodeConflict(err) // v3.1
// 记录创建失败:清理已落盘文件
_ = d.Store.DeleteFile(ctx, savePath)
}
} else {
// 保存失败:尽力清理半写文件
_ = d.Store.DeleteFile(ctx, savePath)
}
releaseStorage(ctx, d.DB, resToken)
if err != nil {
auditUploadEntry(c, "", origName, fh.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件保存失败")
if be, ok := err.(*apiError); ok && be.Status == http.StatusBadRequest {
respondError(c, err) // v3.1:提取码冲突等业务 400 原样透出
} else {
respondError(c, mapStorageError(err))
}
return
}
d.Limiter.Add(c, middleware.LimitUpload)
auditUploadEntry(c, code, origName, fh.Size, fh.Size)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, gin.H{"code": code, "name": origName})
}
// ============ GET/POST /share/metadata 分享元信息(不消耗次数)============
// shareMetadataGET 查询分享元信息(对齐参考 get_file_metadata / build_file_metadata)。
func (d *Deps) shareMetadata(c *gin.Context) {
d.metadataCommon(c, c.Query("code"))
}
// shareMetadataPost JSON 体查询分享元信息(对齐参考 post_file_metadata)。
func (d *Deps) shareMetadataPost(c *gin.Context) {
var body struct {
Code string `json:"code"`
}
if err := bindJSONOrForm(c, &body); err != nil {
auditRecordFailed(c, d.AuditSvc, "请求体格式错误")
respondError(c, err)
return
}
d.metadataCommon(c, body.Code)
}
// metadataCommon 元信息查询公共实现。
func (d *Deps) metadataCommon(c *gin.Context, code string) {
fc, err := d.lookupByCode(c, code, true)
if err != nil {
auditUploadEntry(c, strings.TrimSpace(code), "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, 0)
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, buildFileMetadata(fc))
}
// buildFileMetadata 构造分享元信息(对齐参考 build_file_metadata,不暴露存储路径)。
func buildFileMetadata(fc *model.FileCodes) gin.H {
isText := fc.Text != nil
var remaining any
if fc.ExpiredCount > 0 {
remaining = fc.ExpiredCount
}
var expiredAt any
if fc.ExpiredAt != nil {
expiredAt = fc.ExpiredAt.Format(time.RFC3339)
}
return gin.H{
"code": fc.Code,
"name": fc.Prefix + fc.Suffix,
"size": fc.Size,
"type": map[bool]string{true: "text", false: "file"}[isText],
"is_text": isText,
"created_at": fc.CreatedAt.Format(time.RFC3339),
"expired_at": expiredAt,
"expires_at": expiredAt,
"expired_count": fc.ExpiredCount,
"used_count": fc.UsedCount,
"remaining_downloads": remaining,
}
}
// ============ GET /share/select 取件(消耗次数,流式下载)============
// shareSelect 取件:文本直接返回纯文本;文件流式返回(支持 Range,对齐参考 get_code_file)。
func (d *Deps) shareSelect(c *gin.Context) {
// error 类限流:进入即检查(对齐参考 Depends(ip_limit["error"])
if allowed, _ := d.Limiter.Check(c, middleware.LimitError); !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return
}
code := strings.TrimSpace(c.Query("code"))
fc, err := d.lookupByCode(c, code, true)
if err != nil {
d.Limiter.Add(c, middleware.LimitError) // 取件失败计数
auditUploadEntry(c, code, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
if !d.consumeUsage(c, fc) {
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件已过期")
response.Fail(c, http.StatusNotFound, "文件已过期")
return
}
if fc.Text != nil {
// 文本分享:text/plain 响应(对齐参考 Response(content=text, media_type=text/plain)
name := fc.Prefix + suffixOrTxt(fc)
auditUploadEntry(c, fc.Code, name, fc.Size, int64(len(*fc.Text)))
auditRecordSuccess(c, d.AuditSvc)
c.Header("Content-Disposition", contentDisposition(name))
c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(*fc.Text))
return
}
d.serveFile(c, fc)
}
func suffixOrTxt(fc *model.FileCodes) string {
if fc.Suffix != "" {
return fc.Suffix
}
return ".txt"
}
// ============ POST /share/select 取件详情(JSON,对齐参考 select_file============
// shareSelectPost 返回分享详情 JSON(元信息+内容/下载地址,对齐参考 select_file)。
// 有次数限制的文件返回代理下载地址(消耗发生在 download 时),其余在本次消耗。
func (d *Deps) shareSelectPost(c *gin.Context) {
if allowed, _ := d.Limiter.Check(c, middleware.LimitError); !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return
}
var body struct {
Code string `json:"code"`
}
if err := bindJSONOrForm(c, &body); err != nil {
respondError(c, err)
return
}
fc, err := d.lookupByCode(c, body.Code, true)
if err != nil {
d.Limiter.Add(c, middleware.LimitError)
respondError(c, err)
return
}
detail := buildFileMetadata(fc)
var downloadURL string
if fc.Text != nil {
detail["text"] = *fc.Text
detail["content"] = *fc.Text
} else if fc.ExpiredCount >= 0 {
// 有次数限制:必须经代理下载接口扣次数
downloadURL = d.proxyDownloadURL(fc.Code)
detail["text"] = downloadURL
} else {
// 时间型/永久:优先引擎直链(如 S3 预签名),不支持则代理
if url, err := d.Store.PresignGetURL(c.Request.Context(), fileSavePath(fc), 3600); err == nil {
downloadURL = url
} else {
downloadURL = d.proxyDownloadURL(fc.Code)
}
detail["text"] = downloadURL
}
detail["download_url"] = nil
if downloadURL != "" {
detail["download_url"] = downloadURL
}
// 仅当下载地址不是代理地址时在本次消耗次数(对齐参考 consumes_on_download 判定)
consumesOnDownload := strings.HasPrefix(downloadURL, "/share/download?")
if !consumesOnDownload {
if !d.consumeUsage(c, fc) {
response.Fail(c, http.StatusNotFound, "文件已过期")
return
}
for k, v := range buildFileMetadata(fc) {
detail[k] = v
}
}
response.OK(c, detail)
}
// proxyDownloadURL 生成代理下载地址(对齐参考 get_file_url)。
func (d *Deps) proxyDownloadURL(code string) string {
secret := d.jwtSecret()
return "/share/download?key=" + GetSelectToken(code, secret, 0) + "&code=" + code
}
// ============ GET /share/download 代理下载(token 鉴权,消耗次数)============
// shareDownload 代理下载(对齐参考 download_file):
// key 为 GetSelectToken 生成的窗口令牌,同时接受当前与上一窗口。
func (d *Deps) shareDownload(c *gin.Context) {
key := c.Query("key")
code := strings.TrimSpace(c.Query("code"))
secret := d.jwtSecret()
// L2:HMAC 令牌 + 常量时间比较
if key == "" || secret == "" || !VerifySelectToken(code, secret, key) {
d.Limiter.Add(c, middleware.LimitError)
auditUploadEntry(c, code, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, "下载鉴权失败")
response.Fail(c, http.StatusForbidden, "下载鉴权失败")
return
}
fc, err := d.lookupByCode(c, code, true)
if err != nil {
auditUploadEntry(c, code, "", 0, 0)
auditRecordFailed(c, d.AuditSvc, err.Error())
respondError(c, err)
return
}
if !d.consumeUsage(c, fc) {
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件已过期")
response.Fail(c, http.StatusNotFound, "文件已过期")
return
}
if fc.Text != nil {
auditUploadEntry(c, fc.Code, fc.Prefix+fc.Suffix, fc.Size, int64(len(*fc.Text)))
auditRecordSuccess(c, d.AuditSvc)
response.OK(c, *fc.Text)
return
}
d.serveFile(c, fc)
}
// ============ 文件流式下载(含 Range============
// parseRangeHeader 解析 Range 头(仅 bytes 单区间,多区间/单位错误返回 nil 交由全量处理;
// 合法但越界由引擎返回 ErrRangeNotSatisfiable→416)。
func parseRangeHeader(header string, totalSize int64) *storage.Range {
header = strings.TrimSpace(header)
if header == "" || !strings.HasPrefix(header, "bytes=") {
return nil
}
spec := strings.TrimPrefix(header, "bytes=")
if strings.Contains(spec, ",") { // 多区间不支持,回退全量
return nil
}
dash := strings.Index(spec, "-")
if dash < 0 {
return nil
}
startStr, endStr := strings.TrimSpace(spec[:dash]), strings.TrimSpace(spec[dash+1:])
if startStr == "" {
// bytes=-N:后缀区间
n, err := strconv.ParseInt(endStr, 10, 64)
if err != nil || n <= 0 {
return nil
}
if totalSize > 0 && n >= totalSize {
n = totalSize
}
return &storage.Range{Start: totalSize - n, End: -1}
}
start, err := strconv.ParseInt(startStr, 10, 64)
if err != nil || start < 0 {
return nil
}
if endStr == "" {
return &storage.Range{Start: start, End: -1}
}
end, err := strconv.ParseInt(endStr, 10, 64)
if err != nil || end < start {
return nil
}
return &storage.Range{Start: start, End: end}
}
// serveFile 打开存储流并写响应(200 全量 / 206 区间,自动 Accept-Ranges/Content-Range)。
// 响应字节数由审计中间件包装的 Writer 自动统计。
func (d *Deps) serveFile(c *gin.Context, fc *model.FileCodes) {
ctx := c.Request.Context()
savePath := fileSavePath(fc)
name := fc.Prefix + fc.Suffix
// v3:按文件归属引擎取回(切换引擎后旧文件仍可下载);空戳=历史数据回落当前引擎
store, err := d.storeFor(fc.Engine)
if err != nil {
auditUploadEntry(c, fc.Code, name, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "存储引擎不可用: "+err.Error())
respondError(c, mapStorageError(err))
return
}
// 先 Stat 拿总大小(用于审计与 Range 后缀解析)
var total int64 = -1
if meta, err := store.Stat(ctx, savePath); err == nil && meta != nil {
total = meta.Size
}
rng := parseRangeHeader(c.GetHeader("Range"), total)
dl, err := store.Open(ctx, savePath, rng)
if err != nil {
auditUploadEntry(c, fc.Code, name, fc.Size, 0)
auditRecordFailed(c, d.AuditSvc, "文件读取失败")
respondError(c, mapStorageError(err))
return
}
defer func() { _ = dl.Close() }()
if dl.Total >= 0 {
total = dl.Total
}
c.Header("Accept-Ranges", "bytes")
c.Header("Content-Disposition", contentDisposition(name))
c.Header("Content-Type", "application/octet-stream")
status := http.StatusOK
if rng != nil {
end := dl.End
if end < 0 && total >= 0 {
end = total - 1
}
if end < dl.Start {
auditUploadEntry(c, fc.Code, name, total, 0)
auditRecordFailed(c, d.AuditSvc, "请求范围超出文件大小")
c.Header("Content-Range", fmt.Sprintf("bytes */%d", total))
response.Fail(c, http.StatusRequestedRangeNotSatisfiable, "请求范围超出文件大小")
return
}
status = http.StatusPartialContent
c.Header("Content-Range", fmt.Sprintf("bytes %d-%d/%d", dl.Start, end, total))
c.Header("Content-Length", strconv.FormatInt(end-dl.Start+1, 10))
} else if total >= 0 {
c.Header("Content-Length", strconv.FormatInt(total, 10))
}
auditUploadEntry(c, fc.Code, name, total, 0)
c.Status(status)
n, _ := io.Copy(c.Writer, dl)
middleware.AuditSet(c, func(e *audit.Entry) {
if e.TransferredBytes == 0 {
e.TransferredBytes = n
}
if e.SizeBytes == 0 {
e.SizeBytes = total
}
})
auditRecordSuccess(c, d.AuditSvc)
}
// formInt 读取表单整数(缺省 default 值,非法值亦回退 default)。
func formInt(c *gin.Context, key string, def int) int {
raw := c.PostForm(key)
if raw == "" {
raw = c.Query(key)
}
if raw == "" {
return def
}
n, err := strconv.Atoi(raw)
if err != nil {
return def
}
return n
}
+170
View File
@@ -0,0 +1,170 @@
package api
// v3.1:自定义提取码与站点域名单元测试。
import (
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"github.com/gin-gonic/gin"
)
// postForm 以 urlencoded 表单调用 handlerv3.1 测试辅助)。
func postForm(d *Deps, path string, fields map[string]string) *httptest.ResponseRecorder {
form := url.Values{}
for k, v := range fields {
form.Set(k, v)
}
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
var handler gin.HandlerFunc
switch path {
case "/share/text":
handler = d.shareText
default:
handler = func(c *gin.Context) { c.AbortWithStatus(http.StatusNotFound) }
}
return invoke(handler, req)
}
func TestValidatePickupCode(t *testing.T) {
// 合法:空(用随机码)
if err := validatePickupCode(""); err != nil {
t.Fatalf("空码应合法: %v", err)
}
// 合法:5-8 位字母数字(L3:最小长度由 4 提升至 5)
for _, c := range []string{"abcde", "AB123", "12345678", "a1B2c"} {
if err := validatePickupCode(c); err != nil {
t.Fatalf("合法码 %s 不应报错: %v", c, err)
}
}
// 非法:长度(4 位及以下不再允许)
for _, c := range []string{"abcd", "a1B2", "abc", "123456789"} {
if err := validatePickupCode(c); err == nil {
t.Fatalf("非法长度 %s 应报错", c)
}
}
// 非法:字符
for _, c := range []string{"ab c1", "提码", "ab-cd", "ab.cd", "ab+cd"} {
if err := validatePickupCode(c); err == nil {
t.Fatalf("非法字符 %s 应报错", c)
}
}
}
func TestNormalizeSiteDomain(t *testing.T) {
// 空 = 当前地址
d, err := normalizeSiteDomain("")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "" {
t.Fatalf("want empty, got %q", d)
}
// 完整 URL
d, err = normalizeSiteDomain("https://share.example.com")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "https://share.example.com" {
t.Fatalf("want https://share.example.com, got %q", d)
}
// 带端口 + 去尾斜杠
d, err = normalizeSiteDomain("http://192.168.1.5:8466/")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "http://192.168.1.5:8466" {
t.Fatalf("want http://192.168.1.5:8466, got %q", d)
}
// 裸主机自动补 http
d, err = normalizeSiteDomain("share.example.com")
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if d != "http://share.example.com" {
t.Fatalf("want http://share.example.com, got %q", d)
}
// 非法:路径 / 协议
for _, bad := range []string{"https://a.com/path", "ftp://a.com", "javascript:alert(1)"} {
if _, err := normalizeSiteDomain(bad); err == nil {
t.Fatalf("非法域名 %s 应报错", bad)
}
}
}
// TestShareTextTextPlainCompat 复刻真实浏览器请求形态:
// 旧前端 bundle 发 text/plain Content-Type + urlencoded bodyfetch 字符串 body 默认头)。
// 修复前:该形态被静默存成空文本(bug 1)或 400「分享内容不能为空」(bug 2)。
func TestShareTextTextPlainCompat(t *testing.T) {
d := newPolicyTestDeps(t)
body := strings.NewReader("text=111&expire_value=1&expire_style=day&code=")
req := httptest.NewRequest(http.MethodPost, "/share/text", body)
req.Header.Set("Content-Type", "text/plain;charset=UTF-8")
w := invoke(d.shareText, req)
if w.Code != http.StatusOK {
t.Fatalf("text/plain+urlencoded 应 200: %d %s", w.Code, w.Body.String())
}
// JSON 体但 Content-Type 缺失/为 text/plain 也应可解析
req2 := httptest.NewRequest(http.MethodPost, "/share/text",
strings.NewReader(`{"text":"无头JSON","expire_value":1,"expire_style":"day"}`))
req2.Header.Set("Content-Type", "text/plain;charset=UTF-8")
w2 := invoke(d.shareText, req2)
if w2.Code != http.StatusOK {
t.Fatalf("text/plain+JSON体 应 200: %d %s", w2.Code, w2.Body.String())
}
// 取件确认内容真实落库
req3 := httptest.NewRequest(http.MethodPost, "/share/select",
strings.NewReader(`{"code":"`+codeOf(w)+`"}`))
req3.Header.Set("Content-Type", "application/json")
w3 := invoke(d.shareSelectPost, req3)
if !strings.Contains(w3.Body.String(), "111") {
t.Fatalf("落库内容应为 111: %s", w3.Body.String())
}
}
// codeOf 从创建响应提取取件码。
func codeOf(w *httptest.ResponseRecorder) string {
var env struct {
Data struct {
Code string `json:"code"`
} `json:"data"`
}
_ = json.Unmarshal(w.Body.Bytes(), &env)
return env.Data.Code
}
func TestShareTextCustomCode(t *testing.T) {
d := newPolicyTestDeps(t)
// 自定义码成功创建
w := postForm(d, "/share/text", map[string]string{"text": "自定义码测试", "code": "MYCODE1"})
if w.Code != http.StatusOK {
t.Fatalf("自定义码创建失败: %d %s", w.Code, w.Body.String())
}
// 重复占用 → 400
w = postForm(d, "/share/text", map[string]string{"text": "第二条", "code": "MYCODE1"})
if w.Code != http.StatusBadRequest {
t.Fatalf("占用码应 400: %d %s", w.Code, w.Body.String())
}
// 非法码 → 400
w = postForm(d, "/share/text", map[string]string{"text": "第三条", "code": "abc"})
if w.Code != http.StatusBadRequest {
t.Fatalf("过短码应 400: %d %s", w.Code, w.Body.String())
}
// 空码 → 随机码仍正常
w = postForm(d, "/share/text", map[string]string{"text": "第四条", "code": ""})
if w.Code != http.StatusOK {
t.Fatalf("空码应回退随机: %d %s", w.Code, w.Body.String())
}
}
+75
View File
@@ -0,0 +1,75 @@
package api
import (
"io/fs"
"net/http"
"path"
"strings"
"github.com/gin-gonic/gin"
"filecodebox/internal/response"
web "filecodebox/web"
)
// registerWeb 注册前端静态资源与 SPA 回退(必须最后注册):
// - 静态资源命中 web/dist 内文件则直接服务(带 Immutable 缓存,html 不缓存);
// - 未命中且为 GET/HEAD 且非 /api 前缀:回退 index.html(前端 history 路由
// /s/:code、/admin/*、/docs、/openapi 由 SPA 接管);
// - /api/* 未命中路由:JSON 404(避免调试时拿到 HTML 掩盖真实错误)。
func registerWeb(r *gin.Engine, d *Deps) {
dist, err := web.Dist()
if err != nil {
return // 嵌入异常时跳过(API 仍可用)
}
fileServer := http.StripPrefix("/", http.FileServer(http.FS(dist)))
indexHTML := readIndexHTML(dist)
r.NoRoute(func(c *gin.Context) {
p := c.Request.URL.Path
// 1. API 未命中:JSON 404
if strings.HasPrefix(p, "/api/") || p == "/api" {
response.Fail(c, http.StatusNotFound, "接口不存在")
return
}
// 2. 静态资源命中:直接服务
if p != "/" {
clean := strings.TrimPrefix(path.Clean(p), "/")
if clean != "" {
if f, err := dist.Open(clean); err == nil {
_ = f.Close()
fileServer.ServeHTTP(c.Writer, c.Request)
return
}
}
}
// 3. SPA 回退:仅 GET/HEAD 且接受 HTML
if c.Request.Method == http.MethodGet || c.Request.Method == http.MethodHead {
accept := c.GetHeader("Accept")
if accept == "" || strings.Contains(accept, "text/html") || strings.Contains(accept, "*/*") {
if indexHTML != nil {
c.Data(http.StatusOK, "text/html; charset=utf-8", indexHTML)
return
}
}
// 非 HTML 请求未命中:普通 404
c.Status(http.StatusNotFound)
return
}
c.Status(http.StatusNotFound)
})
}
// readIndexHTML 读取嵌入的 index.htmlSPA 回退用)。
func readIndexHTML(dist fs.FS) []byte {
f, err := dist.Open("index.html")
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
data, err := fs.ReadFile(dist, "index.html")
if err != nil {
return nil
}
return data
}
+212
View File
@@ -0,0 +1,212 @@
// Package audit 提供上传/下载审计日志服务(需求 ③):
// 记录操作时间/IP/UA/设备解析/动作/结果/字节数/耗时,落库 Postgres。
package audit
import (
"context"
"errors"
"log"
"strings"
"sync"
"time"
"gorm.io/gorm"
"filecodebox/internal/model"
)
// Service 审计日志服务。
type Service struct {
sink Sink
}
// Sink 审计落库抽象(生产为 Postgres,测试为内存实现)。
type Sink interface {
// Save 批量落库。
Save(ctx context.Context, logs []model.AuditLog) error
}
// DBSink 基于 GORM 的落库实现。
type DBSink struct{ db *gorm.DB }
// NewDBSink 构造数据库落库实现。
func NewDBSink(db *gorm.DB) *DBSink { return &DBSink{db: db} }
// Save 批量插入审计记录。
func (s *DBSink) Save(ctx context.Context, logs []model.AuditLog) error {
if len(logs) == 0 {
return nil
}
return s.db.WithContext(ctx).CreateInBatches(&logs, 200).Error
}
// NewService 构造审计服务。
func NewService(sink Sink) *Service {
return &Service{sink: sink}
}
// Entry 一次待落库的审计事件。
type Entry struct {
Action string // upload | download
FileCode string // 取件码
FileName string // 原始文件名
SizeBytes int64 // 文件总字节数
TransferredBytes int64 // 实际传输字节数
IP string // 客户端 IP
UserAgent string // User-Agent
DeviceOS string // 操作系统
DeviceBrowser string // 浏览器
DeviceType string // desktop/mobile/tablet/bot/other
Actor string // admin | guest
Result string // success | denied | failed
ErrorMsg string // 失败原因
Duration time.Duration // 耗时
}
// Record 异步写入一条审计日志:先尝试同步落库,失败时进入内存缓冲等待重试,
// 避免审计失败影响主请求,也避免高峰期阻塞。
func (s *Service) Record(entry Entry) {
record := model.AuditLog{
Action: entry.Action,
FileCode: truncate(entry.FileCode, 64),
FileName: truncate(entry.FileName, 255),
SizeBytes: entry.SizeBytes,
TransferredBytes: entry.TransferredBytes,
IP: truncate(entry.IP, 64),
UserAgent: truncate(entry.UserAgent, 512),
DeviceOS: truncate(entry.DeviceOS, 64),
DeviceBrowser: truncate(entry.DeviceBrowser, 64),
DeviceType: truncate(entry.DeviceType, 32),
Actor: truncate(entry.Actor, 64),
Result: normalizeResult(entry.Result),
ErrorMsg: truncate(entry.ErrorMsg, 512),
DurationMs: entry.Duration.Milliseconds(),
}
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := s.sink.Save(ctx, []model.AuditLog{record}); err != nil {
log.Printf("[audit] 审计日志落库失败,进入重试队列: %v", err)
s.enqueue(record)
}
}()
}
// retryBuf 落库失败时的内存重试缓冲。
var retryBuf struct {
sync.Mutex
items []model.AuditLog
}
const maxRetryBuffer = 10000
// enqueue 入队;超出上限时丢弃最旧的,防止内存无限增长。
func (s *Service) enqueue(record model.AuditLog) {
retryBuf.Lock()
if len(retryBuf.items) >= maxRetryBuffer {
retryBuf.items = retryBuf.items[1:]
}
retryBuf.items = append(retryBuf.items, record)
retryBuf.Unlock()
}
// FlushRetry 将缓冲中的审计日志重新落库;由后台定时任务调用。
func (s *Service) FlushRetry() {
retryBuf.Lock()
if len(retryBuf.items) == 0 {
retryBuf.Unlock()
return
}
items := retryBuf.items
retryBuf.items = nil
retryBuf.Unlock()
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := s.sink.Save(ctx, items); err != nil {
log.Printf("[audit] 重试队列落库失败: %v", err)
// 失败则放回队首
retryBuf.Lock()
retryBuf.items = append(items, retryBuf.items...)
if len(retryBuf.items) > maxRetryBuffer {
retryBuf.items = retryBuf.items[:maxRetryBuffer]
}
retryBuf.Unlock()
}
}
// StartRetryLoop 启动后台重试循环。
func (s *Service) StartRetryLoop(stop <-chan struct{}) {
go func() {
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-stop:
s.FlushRetry()
return
case <-ticker.C:
s.FlushRetry()
}
}
}()
}
// Query 按条件分页查询审计日志(管理端使用)。
// action/ip/result 为可选过滤;begin/end 为创建时间范围(可选)。
func (s *Service) Query(page, pageSize int, action, ip, result string, begin, end *time.Time) ([]model.AuditLog, int64, error) {
dbSink, ok := s.sink.(*DBSink)
if !ok {
return nil, 0, errors.New("audit: 当前 sink 不支持查询")
}
if page < 1 {
page = 1
}
if pageSize < 1 || pageSize > 200 {
pageSize = 20
}
q := dbSink.db.Model(&model.AuditLog{})
if action != "" {
q = q.Where("action = ?", action)
}
if ip != "" {
q = q.Where("ip = ?", ip)
}
if result != "" {
q = q.Where("result = ?", result)
}
if begin != nil {
q = q.Where("created_at >= ?", *begin)
}
if end != nil {
q = q.Where("created_at <= ?", *end)
}
var total int64
if err := q.Count(&total).Error; err != nil {
return nil, 0, err
}
var logs []model.AuditLog
err := q.Order("id DESC").
Offset((page - 1) * pageSize).
Limit(pageSize).
Find(&logs).Error
return logs, total, err
}
func truncate(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n]
}
func normalizeResult(r string) string {
switch strings.TrimSpace(r) {
case model.AuditResultSuccess, model.AuditResultDenied, model.AuditResultFailed:
return strings.TrimSpace(r)
case "":
return model.AuditResultFailed
default:
return model.AuditResultFailed
}
}
+84
View File
@@ -0,0 +1,84 @@
package audit
import (
"strings"
)
// DeviceInfo 从 User-Agent 解析出的设备信息。
type DeviceInfo struct {
OS string // Windows/macOS/Android/iOS/Linux/Unknown
Browser string // Chrome/Firefox/Safari/Edge/Other
Type string // desktop/mobile/tablet/bot/other
}
// 动作常量。
const (
ActionUpload = "upload"
ActionDownload = "download"
// ActionAdmin 管理端敏感操作(登录/登出/配置/密码/引擎切换/文件删除等),
// L5:纳入审计以便追溯登录失败与配置变更。
ActionAdmin = "admin"
)
// 角色常量。
const (
ActorAdmin = "admin"
ActorGuest = "guest"
)
// botKeywords 常见爬虫/机器人标识。
var botKeywords = []string{"bot", "spider", "crawl", "slurp", "curl/", "wget", "python-requests", "go-http-client"}
// ParseUserAgent 解析 User-Agent 为设备信息(轻量规则,避免引入重依赖)。
func ParseUserAgent(ua string) DeviceInfo {
ua = strings.TrimSpace(ua)
lower := strings.ToLower(ua)
if ua == "" {
return DeviceInfo{OS: "Unknown", Browser: "Other", Type: "other"}
}
for _, kw := range botKeywords {
if strings.Contains(lower, kw) {
return DeviceInfo{OS: "Unknown", Browser: "Other", Type: "bot"}
}
}
info := DeviceInfo{OS: "Unknown", Browser: "Other", Type: "desktop"}
// 操作系统
switch {
case strings.Contains(lower, "windows"):
info.OS = "Windows"
case strings.Contains(lower, "iphone"), strings.Contains(lower, "ipod"):
info.OS = "iOS"
info.Type = "mobile"
case strings.Contains(lower, "ipad"):
info.OS = "iOS"
info.Type = "tablet"
case strings.Contains(lower, "mac os x"), strings.Contains(lower, "macintosh"):
info.OS = "macOS"
case strings.Contains(lower, "android"):
info.OS = "Android"
info.Type = "mobile"
if strings.Contains(lower, "tablet") || !strings.Contains(lower, "mobile") {
info.Type = "tablet"
}
case strings.Contains(lower, "linux"), strings.Contains(lower, "ubuntu"), strings.Contains(lower, "fedora"):
info.OS = "Linux"
}
// 浏览器(顺序重要:Edge/OPR 必须在 Chrome 之前判断)
switch {
case strings.Contains(lower, "edg/"), strings.Contains(lower, "edge/"):
info.Browser = "Edge"
case strings.Contains(lower, "opr/"), strings.Contains(lower, "opera"):
info.Browser = "Opera"
case strings.Contains(lower, "chrome/"), strings.Contains(lower, "crios/"):
info.Browser = "Chrome"
case strings.Contains(lower, "firefox/"), strings.Contains(lower, "fxios/"):
info.Browser = "Firefox"
case strings.Contains(lower, "safari/"):
info.Browser = "Safari"
}
return info
}
+45
View File
@@ -0,0 +1,45 @@
package audit
import "testing"
func TestParseUserAgent(t *testing.T) {
cases := []struct {
ua string
os string
browser string
typ string
}{
{
ua: "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36",
os: "Windows", browser: "Chrome", typ: "desktop",
},
{
ua: "Mozilla/5.0 (iPhone; CPU iPhone OS 17_0 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.0 Mobile/15E148 Safari/604.1",
os: "iOS", browser: "Safari", typ: "mobile",
},
{
ua: "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36 Edg/119.0.0.0",
os: "macOS", browser: "Edge", typ: "desktop",
},
{
ua: "Mozilla/5.0 (Linux; Android 13; Pixel 7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Mobile Safari/537.36",
os: "Android", browser: "Chrome", typ: "mobile",
},
{
ua: "curl/8.4.0",
os: "Unknown", browser: "Other", typ: "bot",
},
{
ua: "Mozilla/5.0 (X11; Linux x86_64; rv:121.0) Gecko/20100101 Firefox/121.0",
os: "Linux", browser: "Firefox", typ: "desktop",
},
{ua: "", os: "Unknown", browser: "Other", typ: "other"},
}
for i, tc := range cases {
got := ParseUserAgent(tc.ua)
if got.OS != tc.os || got.Browser != tc.browser || got.Type != tc.typ {
t.Errorf("case %d: ParseUserAgent(%q) = %+v, want os=%s browser=%s type=%s",
i, tc.ua, got, tc.os, tc.browser, tc.typ)
}
}
}
+42
View File
@@ -0,0 +1,42 @@
// Package cache 提供统一缓存接口:FCB_REDIS_ADDR 未配置时自动降级为进程内存实现,
// 用于 IP 限流计数与热点配置缓存(需求 ② 的可选 Redis 增强)。
package cache
import (
"context"
"errors"
"time"
)
// ErrNotFound 表示键不存在。
var ErrNotFound = errors.New("cache: key 不存在")
// Cache 缓存统一接口。
type Cache interface {
// Get 读取字符串值;键不存在返回 ErrNotFound。
Get(ctx context.Context, key string) (string, error)
// Set 写入字符串值,ttl<=0 表示不过期。
Set(ctx context.Context, key, value string, ttl time.Duration) error
// Delete 删除键。
Delete(ctx context.Context, keys ...string) error
// Exists 判断键是否存在。
Exists(ctx context.Context, key string) (bool, error)
// Incr 原子自增;键不存在时从 0 开始并设置 ttl 窗口(限流固定窗口用)。
Incr(ctx context.Context, key string, ttl time.Duration) (int64, error)
// Close 释放底层资源(Redis 连接;内存实现为空操作)。
Close() error
}
// RedisOptions Redis 连接参数(addr 为空 → 内存实现;db 为 FCB_REDIS_DB 库号)。
type RedisOptions struct {
Addr string
DB int // 逻辑库号 0-15cluster 模式忽略)
}
// New 按配置构造缓存实现:redisAddr 为空 → 内存实现。
func New(ctx context.Context, opt RedisOptions) (Cache, error) {
if opt.Addr == "" {
return NewMemory(), nil
}
return NewRedis(ctx, opt.Addr, opt.DB)
}
+81
View File
@@ -0,0 +1,81 @@
package cache
import (
"context"
"sync"
"testing"
"time"
)
func TestMemoryCacheSetGet(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
if err := c.Set(ctx, "k1", "v1", 0); err != nil {
t.Fatalf("Set 失败: %v", err)
}
v, err := c.Get(ctx, "k1")
if err != nil || v != "v1" {
t.Fatalf("Get = (%q, %v)", v, err)
}
if _, err := c.Get(ctx, "missing"); err != ErrNotFound {
t.Fatalf("缺失键应返回 ErrNotFound: %v", err)
}
_ = c.Delete(ctx, "k1")
if _, err := c.Get(ctx, "k1"); err != ErrNotFound {
t.Fatal("删除后应不存在")
}
}
func TestMemoryCacheTTL(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
_ = c.Set(ctx, "ttl", "x", 50*time.Millisecond)
if ok, _ := c.Exists(ctx, "ttl"); !ok {
t.Fatal("TTL 内应存在")
}
time.Sleep(80 * time.Millisecond)
if _, err := c.Get(ctx, "ttl"); err != ErrNotFound {
t.Fatal("过期后应不存在")
}
}
func TestMemoryCacheIncrWindow(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
for i := int64(1); i <= 3; i++ {
n, err := c.Incr(ctx, "rl", time.Minute)
if err != nil || n != i {
t.Fatalf("Incr = (%d, %v), want (%d, nil)", n, err, i)
}
}
// 窗口过期后重新计数
_ = c.Set(ctx, "short", "seed", time.Millisecond)
time.Sleep(5 * time.Millisecond)
n, err := c.Incr(ctx, "short", time.Millisecond)
if err != nil || n != 1 {
t.Fatalf("过期窗口重置失败: (%d, %v)", n, err)
}
}
func TestMemoryCacheConcurrentIncr(t *testing.T) {
c := NewMemory()
defer c.Close()
ctx := context.Background()
var wg sync.WaitGroup
for i := 0; i < 50; i++ {
wg.Add(1)
go func() {
defer wg.Done()
_, _ = c.Incr(ctx, "cnt", time.Minute)
}()
}
wg.Wait()
n, _ := c.Incr(ctx, "cnt", time.Minute)
if n != 51 {
t.Fatalf("并发计数丢失: %d != 51", n)
}
}
+159
View File
@@ -0,0 +1,159 @@
package cache
import (
"context"
"sync"
"time"
)
// memoryItem 内存缓存条目。
type memoryItem struct {
value string
expiresAt time.Time // 零值表示不过期
}
// MemoryCache 进程内存缓存实现(单机、无持久化)。
type MemoryCache struct {
mu sync.RWMutex
items map[string]memoryItem
done chan struct{}
}
// NewMemory 构造内存缓存,并启动后台过期清理。
func NewMemory() *MemoryCache {
m := &MemoryCache{
items: make(map[string]memoryItem),
done: make(chan struct{}),
}
go m.gcLoop()
return m
}
// gcLoop 每分钟清理一次过期键,避免长期运行内存膨胀。
func (m *MemoryCache) gcLoop() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for {
select {
case <-m.done:
return
case now := <-ticker.C:
m.mu.Lock()
for k, item := range m.items {
if !item.expiresAt.IsZero() && now.After(item.expiresAt) {
delete(m.items, k)
}
}
m.mu.Unlock()
}
}
}
// Get 读取键值。
func (m *MemoryCache) Get(_ context.Context, key string) (string, error) {
m.mu.RLock()
item, ok := m.items[key]
m.mu.RUnlock()
if !ok {
return "", ErrNotFound
}
if !item.expiresAt.IsZero() && time.Now().After(item.expiresAt) {
return "", ErrNotFound
}
return item.value, nil
}
// Set 写入键值。
func (m *MemoryCache) Set(_ context.Context, key, value string, ttl time.Duration) error {
item := memoryItem{value: value}
if ttl > 0 {
item.expiresAt = time.Now().Add(ttl)
}
m.mu.Lock()
m.items[key] = item
m.mu.Unlock()
return nil
}
// Delete 删除键。
func (m *MemoryCache) Delete(_ context.Context, keys ...string) error {
m.mu.Lock()
for _, k := range keys {
delete(m.items, k)
}
m.mu.Unlock()
return nil
}
// Exists 判断键是否存在。
func (m *MemoryCache) Exists(_ context.Context, key string) (bool, error) {
_, err := m.Get(context.Background(), key)
return err == nil, nil
}
// Incr 原子自增;首次创建时记录窗口起点(以过期时间体现)。
func (m *MemoryCache) Incr(_ context.Context, key string, ttl time.Duration) (int64, error) {
m.mu.Lock()
defer m.mu.Unlock()
now := time.Now()
item, ok := m.items[key]
if ok && !item.expiresAt.IsZero() && now.After(item.expiresAt) {
// 窗口已过期,重新计数
ok = false
}
var n int64
if !ok {
n = 1
newItem := memoryItem{value: "1"}
if ttl > 0 {
newItem.expiresAt = now.Add(ttl)
}
m.items[key] = newItem
return n, nil
}
// 解析现有值
for _, c := range item.value {
if c < '0' || c > '9' {
n = 0
break
}
n = n*10 + int64(c-'0')
}
n++
newItem := memoryItem{value: itoa(n), expiresAt: item.expiresAt}
m.items[key] = newItem
return n, nil
}
// Close 停止清理协程。
func (m *MemoryCache) Close() error {
select {
case <-m.done:
default:
close(m.done)
}
return nil
}
// itoa 简单整数转字符串,避免在锁内依赖 strconv 的额外开销(数值都很小)。
func itoa(n int64) string {
if n == 0 {
return "0"
}
var buf [20]byte
i := len(buf)
neg := n < 0
if neg {
n = -n
}
for n > 0 {
i--
buf[i] = byte('0' + n%10)
n /= 10
}
if neg {
i--
buf[i] = '-'
}
return string(buf[i:])
}
+117
View File
@@ -0,0 +1,117 @@
package cache
import (
"context"
"fmt"
"net/url"
"strings"
"time"
"github.com/redis/go-redis/v9"
)
// RedisCache 基于 Redis 的缓存实现(可选增强)。
type RedisCache struct {
client *redis.Client
}
// NewRedis 连接 Redis 并校验可用性。addr 支持两种形式:
// - host:port(纯地址,库号由 db 参数指定)
// - redis://[:password@]host:port[/db]URL 形式,URL 中的库号优先于 db 参数)
func NewRedis(ctx context.Context, addr string, db int) (*RedisCache, error) {
opts, err := buildRedisOptions(addr, db)
if err != nil {
return nil, err
}
client := redis.NewClient(opts)
if err := client.Ping(ctx).Err(); err != nil {
_ = client.Close()
return nil, fmt.Errorf("cache: Redis 连接失败 %s: %w", addr, err)
}
return &RedisCache{client: client}, nil
}
// buildRedisOptions 构造 go-redis 连接选项(纯地址 / URL 形式统一入口)。
func buildRedisOptions(addr string, db int) (*redis.Options, error) {
opts := &redis.Options{
Addr: addr,
DB: db,
DialTimeout: 5 * time.Second,
ReadTimeout: 3 * time.Second,
WriteTimeout: 3 * time.Second,
PoolSize: 32,
}
if strings.HasPrefix(addr, "redis://") || strings.HasPrefix(addr, "rediss://") {
u, err := redis.ParseURL(addr)
if err != nil {
return nil, fmt.Errorf("cache: Redis 地址解析失败 %s: %w", addr, err)
}
// URL 未显式携带库号(路径为空或 /)时用 db 参数;显式 /N 优先
if u.DB == 0 && !urlHasDBPath(addr) {
u.DB = db
}
u.DialTimeout = opts.DialTimeout
u.ReadTimeout = opts.ReadTimeout
u.WriteTimeout = opts.WriteTimeout
u.PoolSize = opts.PoolSize
opts = u
}
return opts, nil
}
// urlHasDBPath 判断 redis:// URL 是否显式携带了库号路径(如 /5)。
func urlHasDBPath(raw string) bool {
u, err := url.Parse(raw)
if err != nil {
return false
}
return strings.Trim(u.Path, "/") != ""
}
// Get 读取键值。
func (r *RedisCache) Get(ctx context.Context, key string) (string, error) {
val, err := r.client.Get(ctx, key).Result()
if err == redis.Nil {
return "", ErrNotFound
}
return val, err
}
// Set 写入键值。
func (r *RedisCache) Set(ctx context.Context, key, value string, ttl time.Duration) error {
return r.client.Set(ctx, key, value, ttl).Err()
}
// Delete 删除键。
func (r *RedisCache) Delete(ctx context.Context, keys ...string) error {
if len(keys) == 0 {
return nil
}
return r.client.Del(ctx, keys...).Err()
}
// Exists 判断键是否存在。
func (r *RedisCache) Exists(ctx context.Context, key string) (bool, error) {
n, err := r.client.Exists(ctx, key).Result()
return n > 0, err
}
// Incr 原子自增;首次创建时设置窗口 TTL。
// 使用 Lua 脚本保证 INCR+EXPIRE 原子性,避免多实例下窗口被反复重置。
func (r *RedisCache) Incr(ctx context.Context, key string, ttl time.Duration) (int64, error) {
var incrScript = redis.NewScript(`
local n = redis.call('INCR', KEYS[1])
if n == 1 and ARGV[1] ~= '0' then
redis.call('PEXPIRE', KEYS[1], ARGV[1])
end
return n
`)
ttlMs := int64(0)
if ttl > 0 {
ttlMs = ttl.Milliseconds()
}
return incrScript.Run(ctx, r.client, []string{key}, ttlMs).Int64()
}
// Close 关闭 Redis 连接。
func (r *RedisCache) Close() error { return r.client.Close() }
+70
View File
@@ -0,0 +1,70 @@
// redis_options_test.go — FCB_REDIS_DB / URL 库号解析单测。
package cache
import "testing"
func TestBuildRedisOptionsPlainAddr(t *testing.T) {
opts, err := buildRedisOptions("127.0.0.1:6379", 0)
if err != nil {
t.Fatalf("plain addr: %v", err)
}
if opts.DB != 0 {
t.Fatalf("默认库号应为 0, got %d", opts.DB)
}
opts, err = buildRedisOptions("127.0.0.1:6379", 5)
if err != nil {
t.Fatalf("plain addr db=5: %v", err)
}
if opts.Addr != "127.0.0.1:6379" || opts.DB != 5 {
t.Fatalf("host:port + db: got addr=%s db=%d", opts.Addr, opts.DB)
}
}
func TestBuildRedisOptionsURL(t *testing.T) {
cases := []struct {
name string
url string
dbParam int
wantDB int
wantPw string
}{
{"URL 无库号用参数", "redis://127.0.0.1:6379", 3, 3, ""},
{"URL 显式库号优先", "redis://127.0.0.1:6379/7", 3, 7, ""},
{"URL 带密码", "redis://:secretpw@127.0.0.1:6379/2", 0, 2, "secretpw"},
{"rediss 无库号用参数", "rediss://127.0.0.1:6379", 9, 9, ""},
{"URL 根路径视为无库号", "redis://127.0.0.1:6379/", 4, 4, ""},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
opts, err := buildRedisOptions(tc.url, tc.dbParam)
if err != nil {
t.Fatalf("buildRedisOptions(%q): %v", tc.url, err)
}
if opts.DB != tc.wantDB {
t.Fatalf("db = %d, want %d", opts.DB, tc.wantDB)
}
if opts.Password != tc.wantPw {
t.Fatalf("password = %q, want %q", opts.Password, tc.wantPw)
}
if opts.Addr != "127.0.0.1:6379" {
t.Fatalf("addr = %q", opts.Addr)
}
})
}
}
func TestBuildRedisOptionsInvalidURL(t *testing.T) {
if _, err := buildRedisOptions("redis://[bad", 0); err == nil {
t.Fatal("非法 URL 应报错")
}
}
func TestURLHasDBPath(t *testing.T) {
if urlHasDBPath("redis://h:6379") || urlHasDBPath("redis://h:6379/") {
t.Fatal("无路径或根路径应视为 false")
}
if !urlHasDBPath("redis://h:6379/5") {
t.Fatal("/5 应视为 true")
}
}
+417
View File
@@ -0,0 +1,417 @@
// Package config 提供全局配置:默认值对齐参考实现 core/settings.py
// 支持 FCB_* 环境变量覆盖默认值,再由数据库 settings KV 做运行时覆盖。
package config
import (
"fmt"
"os"
"strconv"
"strings"
)
// 会话有效期边界(与参考实现保持一致:天级、可配 1~365 天)。
const (
// AdminSessionExpireDefault 默认 7 天(L8:由 30 天缩短,降低 localStorage
// token 泄露后的暴露窗口;管理员可在 1~365 天内自行调整)
AdminSessionExpireDefault = 7 * 24 * 60 * 60
AdminSessionExpireMin = 24 * 60 * 60 // 最小 1 天
AdminSessionExpireMax = 365 * 24 * 60 * 60 // 最大 365 天
// DefaultSQLitePath SQLite 模式默认数据库文件路径(相对运行目录,自动创建 data/)。
DefaultSQLitePath = "./data/filecodebox.db"
)
// 数据库驱动常量(需求 ⑧:SQLite 默认、Postgres 可选)。
const (
DBDriverSQLite = "sqlite"
DBDriverPostgres = "postgres"
)
// DefaultLogoURL / DefaultFaviconURL 默认 Logo 与 favicon(需求 ⑤):
// v2 起默认改用前端打包的本地资源(web/src/assets/brand/logo.svg + favicon.png
// 经 Vite 产出 /assets/logo-*.svg 与 /assets/favicon-*.png)。此处留空,
// GET /api/v1/config 下发空值时前端 displayLogoUrl/displayFaviconUrl 回落到本地打包资源;
// 管理端仍可设置任意 URL 全站替换。
const DefaultLogoURL = ""
// DefaultFaviconURL favicon/备用 Logo 默认空串(语义见 DefaultLogoURL 注释)。
const DefaultFaviconURL = ""
// Config 运行时配置。Env 为 FCB_* 环境变量解析结果(进程级),
// KV 为数据库 settings 键值覆盖(可被管理端动态修改)。
type Config struct {
Env *EnvConfig
KV map[string]any
}
// EnvConfig 进程级环境变量配置,仅能通过环境变量修改。
type EnvConfig struct {
DBDriver string // FCB_DB_DRIVERsqlite|postgres,默认 sqlite(需求 ⑧)
DBDSN string // FCB_DB_DSNpostgres 必需;sqlite 为空时用 DefaultSQLitePath
RedisAddr string // FCB_REDIS_ADDR,可选;为空时缓存降级为内存实现
RedisDB int // FCB_REDIS_DBRedis 逻辑库号 0-15,默认 0(URL 形式地址以 URL 内库号优先)
Listen string // FCB_LISTEN,监听地址,默认 :8466
StorageEngine string // FCB_STORAGE_ENGINElocal|s3|webdav,默认 local
TrustedProxies []string // FCB_TRUSTED_PROXIES,逗号分隔的可信代理 CIDR
}
// defaults 返回与参考实现 core/settings.py DEFAULT_CONFIG 对齐的默认配置。
func defaults() map[string]any {
return map[string]any{
// 存储引擎与路径
"file_storage": "local",
"storage_path": "",
"storageLimit": 0,
// v3:存储引擎运行时可配(热切换);空=沿用 Env.StorageEngine 启动值
"storage_engine": "",
"site_domain": "",
// 站点信息
"name": "文件快传",
"site_name": "文件快传", // 新增:管理端可自定义
"description": "开箱即用的文件快传系统",
"notify_title": "系统通知",
"notify_content": "欢迎使用文件快传,拖拽或粘贴即可分享文本与文件。",
"page_explain": "请勿上传或分享违法内容。根据《中华人民共和国网络安全法》、《中华人民共和国刑法》、《中华人民共和国治安管理处罚法》等相关规定。 传播或存储违法、违规内容,会受到相关处罚,严重者将承担刑事责任。本站坚决配合相关部门,确保网络内容的安全,和谐,打造绿色网络环境。",
"keywords": "文件快传, 文件分享, 匿名口令分享文本, 文件",
// 需求 ⑤:默认 Logo 与 favicon(空 = 前端使用打包的本地资源)
"logo_url": DefaultLogoURL,
"favicon_url": DefaultFaviconURL,
// 需求 ①:背景图(v2 新增 background_urlbackground 为参考实现既有键,保留兼容)
"background": "",
"background_url": "",
// 需求 ②:页脚自定义内容与备案号
"footer_text": "",
"footer_beian": "",
// 需求 ③:系统通知(notify_enabled 新增开关,title/content 沿用参考语义)
"notify_enabled": 1,
// 需求 ④:保存策略(次数上限新增;时间上限沿用 max_save_seconds
"max_save_count": 0,
// 需求 ⑩:存储策略-单文件上限(0=回落 uploadSize,避免与参考键冲突)
"max_file_size": 0,
// 本地存储
"local_storage_path": "",
// S3 引擎
"s3_access_key_id": "",
"s3_secret_access_key": "",
"s3_bucket_name": "",
"s3_endpoint_url": "",
"s3_region_name": "auto",
"s3_signature_version": "s3v4",
"s3_hostname": "",
"s3_addressing_style": "auto",
"s3_proxy": 0,
"aws_session_token": "",
// WebDAV 引擎
"webdav_url": "",
"webdav_username": "",
"webdav_password": "",
"webdav_root_path": "filebox_storage",
"webdav_proxy": 0,
// 安全
"admin_token": "", // 管理员密码哈希;为空表示未初始化
"jwt_secret": "",
"adminSessionExpire": AdminSessionExpireDefault,
// 上传与分享策略
"openUpload": 1,
"uploadSize": 1024 * 1024 * 10,
"allowed_file_types": []string{"*"},
"expireStyle": []string{"day", "hour", "minute", "forever", "count"},
"code_generate_type": "secret",
"uploadMinute": 1,
"uploadCount": 10,
"errorMinute": 1,
"errorCount": 10,
"loginCount": 5,
"loginMinute": 15,
"max_save_seconds": 0,
"enableChunk": 0,
// 界面
"opacity": 0.9,
"showAdminAddr": 0,
"robotsText": "User-agent: *\nDisallow: /",
"serverWorkers": 1,
"serverHost": "0.0.0.0",
"serverPort": 8466,
}
}
// loadEnv 解析 FCB_* 环境变量;返回 nil 表示未设置任何必需项。
func loadEnv() (*EnvConfig, error) {
env := &EnvConfig{
DBDriver: strings.ToLower(strings.TrimSpace(os.Getenv("FCB_DB_DRIVER"))),
DBDSN: strings.TrimSpace(os.Getenv("FCB_DB_DSN")),
RedisAddr: strings.TrimSpace(os.Getenv("FCB_REDIS_ADDR")),
Listen: strings.TrimSpace(os.Getenv("FCB_LISTEN")),
StorageEngine: strings.TrimSpace(os.Getenv("FCB_STORAGE_ENGINE")),
}
// Redis 库号(FCB_REDIS_DB0-15;非法值忽略用默认 0)
if v := strings.TrimSpace(os.Getenv("FCB_REDIS_DB")); v != "" {
if n, err := strconv.Atoi(v); err == nil && n >= 0 && n <= 15 {
env.RedisDB = n
}
}
if env.Listen == "" {
env.Listen = ":8466"
}
if env.StorageEngine == "" {
env.StorageEngine = "local"
}
switch env.StorageEngine {
case "local", "s3", "webdav":
default:
return nil, fmt.Errorf("FCB_STORAGE_ENGINE 无效值 %q,仅支持 local|s3|webdav", env.StorageEngine)
}
if raw := strings.TrimSpace(os.Getenv("FCB_TRUSTED_PROXIES")); raw != "" {
for _, item := range strings.Split(raw, ",") {
if item = strings.TrimSpace(item); item != "" {
env.TrustedProxies = append(env.TrustedProxies, item)
}
}
}
return env, nil
}
// New 从环境变量构造配置;KV 覆盖先为空。
// 需求 ⑧:FCB_DB_DRIVER 默认 sqlite(零依赖);postgres 必须提供 FCB_DB_DSN。
func New() (*Config, error) {
env, err := loadEnv()
if err != nil {
return nil, err
}
switch env.DBDriver {
case "", DBDriverSQLite:
env.DBDriver = DBDriverSQLite
// sqlite 模式 DSN 可为空:数据库层回退到 DefaultSQLitePath
case DBDriverPostgres:
if env.DBDSN == "" {
return nil, fmt.Errorf("FCB_DB_DRIVER=postgres 时必须提供 FCB_DB_DSNPostgres 连接串)")
}
default:
return nil, fmt.Errorf("FCB_DB_DRIVER 无效值 %q,仅支持 sqlite|postgres", env.DBDriver)
}
return &Config{Env: env, KV: map[string]any{}}, nil
}
// ApplyKV 用数据库 settings KV 覆盖运行时配置(内部键以 _ 开头的不允许覆盖)。
func (c *Config) ApplyKV(kv map[string]any) {
for k, v := range kv {
if strings.HasPrefix(k, "_") {
continue
}
c.KV[k] = v
}
}
// Get 按 键读取:KV 覆盖 > 默认值;找不到返回零值与 false。
func (c *Config) Get(key string) (any, bool) {
if v, ok := c.KV[key]; ok {
return v, true
}
v, ok := defaults()[key]
return v, ok
}
// GetString 取字符串配置。
func (c *Config) GetString(key string) string {
v, ok := c.Get(key)
if !ok || v == nil {
return ""
}
if s, ok := v.(string); ok {
return s
}
return fmt.Sprintf("%v", v)
}
// GetInt 取整型配置,兼容 JSON 数字(float64)与字符串。
func (c *Config) GetInt(key string) int {
n, _ := c.getInt64(key)
return int(n)
}
// GetInt64 取长整型配置。
func (c *Config) GetInt64(key string) int64 {
n, _ := c.getInt64(key)
return n
}
func (c *Config) getInt64(key string) (int64, bool) {
v, ok := c.Get(key)
if !ok || v == nil {
return 0, false
}
switch n := v.(type) {
case int:
return int64(n), true
case int64:
return n, true
case float64:
return int64(n), true
case string:
if n, err := strconv.ParseInt(strings.TrimSpace(n), 10, 64); err == nil {
return n, true
}
}
return 0, false
}
// GetBool 取布尔配置,兼容 1/0、"true"/"false"/"on"/"yes"。
func (c *Config) GetBool(key string) bool {
v, ok := c.Get(key)
if !ok || v == nil {
return false
}
switch b := v.(type) {
case bool:
return b
case int:
return b != 0
case float64:
return b != 0
case string:
switch strings.ToLower(strings.TrimSpace(b)) {
case "1", "true", "on", "yes":
return true
}
}
return false
}
// GetStringSlice 取字符串切片配置。
// SiteDomain 站点对外域名(v3.1):空=分享链接用当前访问地址。
func (c *Config) SiteDomain() string {
return strings.TrimRight(strings.TrimSpace(c.GetString("site_domain")), "/")
}
func (c *Config) GetStringSlice(key string) []string {
v, ok := c.Get(key)
if !ok || v == nil {
return nil
}
switch s := v.(type) {
case []string:
return s
case []any:
out := make([]string, 0, len(s))
for _, item := range s {
if item == nil {
continue
}
out = append(out, fmt.Sprintf("%v", item))
}
return out
case string:
var out []string
for _, item := range strings.Split(s, ",") {
if item = strings.TrimSpace(item); item != "" {
out = append(out, item)
}
}
return out
}
return nil
}
// —— 常用字段的便捷访问(与参考 settings.xxx 对齐)——
// SiteName 站点名称。
func (c *Config) SiteName() string {
if v := c.GetString("site_name"); v != "" {
return v
}
return c.GetString("name")
}
// LogoURL 页面 Logo。
func (c *Config) LogoURL() string { return c.GetString("logo_url") }
// FaviconURL favicon 地址。
func (c *Config) FaviconURL() string { return c.GetString("favicon_url") }
// OpenUpload 是否允许游客上传。
func (c *Config) OpenUpload() bool { return c.GetBool("openUpload") }
// UploadSize 单文件大小上限(字节)。
func (c *Config) UploadSize() int64 { return c.GetInt64("uploadSize") }
// AllowedFileTypes 允许的文件类型列表("*" 表示不限制)。
func (c *Config) AllowedFileTypes() []string { return c.GetStringSlice("allowed_file_types") }
// ExpireStyle 允许的过期方式。
func (c *Config) ExpireStyle() []string { return c.GetStringSlice("expireStyle") }
// EnableChunk 是否启用分片上传。
func (c *Config) EnableChunk() bool { return c.GetBool("enableChunk") }
// MaxSaveSeconds 最长保存秒数,0 表示不限制。
func (c *Config) MaxSaveSeconds() int64 { return c.GetInt64("max_save_seconds") }
// MaxSaveCount 单次分享最大可取次数上限(需求 ④),0 表示不限制。
func (c *Config) MaxSaveCount() int { return c.GetInt("max_save_count") }
// MaxFileSize 存储策略-单文件上限(需求 ⑩);0 表示回落 uploadSize。
func (c *Config) MaxFileSize() int64 {
if n := c.GetInt64("max_file_size"); n > 0 {
return n
}
return c.UploadSize()
}
// FooterText 页脚自定义内容(需求 ②)。
func (c *Config) FooterText() string { return c.GetString("footer_text") }
// FooterBeian 备案号(需求 ②)。
func (c *Config) FooterBeian() string { return c.GetString("footer_beian") }
// BackgroundURL 背景图地址(需求 ①);空表示使用主题默认。
func (c *Config) BackgroundURL() string {
if v := c.GetString("background_url"); v != "" {
return v
}
return c.GetString("background")
}
// NotifyEnabled 系统通知开关(需求 ③):默认开启。
func (c *Config) NotifyEnabled() bool {
if v, ok := c.Get("notify_enabled"); ok && v != nil {
return c.GetBool("notify_enabled")
}
return true
}
// SQLitePath 数据库文件路径:sqlite 模式下 DSN 为空时回退默认路径(需求 ⑧)。
func (c *Config) SQLitePath() string {
if c.Env.DBDriver != DBDriverSQLite {
return ""
}
if c.Env.DBDSN != "" {
return c.Env.DBDSN
}
return DefaultSQLitePath
}
// AdminSessionExpireSeconds 管理员会话有效期(秒),
// 参考 apps/admin/dependencies.py 的 get_admin_session_expire_seconds。
func (c *Config) AdminSessionExpireSeconds() int {
n := c.GetInt("adminSessionExpire")
if n < AdminSessionExpireMin || n > AdminSessionExpireMax || n%AdminSessionExpireMin != 0 {
return AdminSessionExpireDefault
}
return n
}
// Engine 当前存储引擎。
// Engine 返回当前存储引擎名:KV storage_engine 优先(v3 运行时可改),
// 空(未设置/历史数据)回落启动值 Env.StorageEngineenv 校验过的 local|s3|webdav)。
// 枚举校验内联(避免 config→storage 反向依赖)。
func (c *Config) Engine() string {
if v, ok := c.Get(KeyStorageEngine); ok {
if s, isStr := v.(string); isStr {
switch s {
case "local", "s3", "webdav":
return s
}
}
}
return c.Env.StorageEngine
}
+160
View File
@@ -0,0 +1,160 @@
package config
import (
"testing"
)
func TestNewDefaultsToSQLite(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "")
t.Setenv("FCB_DB_DSN", "")
t.Setenv("FCB_REDIS_ADDR", "")
t.Setenv("FCB_LISTEN", "")
t.Setenv("FCB_STORAGE_ENGINE", "")
c, err := New()
if err != nil {
t.Fatalf("默认(无 DSN)应可构造: %v", err)
}
if c.Env.DBDriver != DBDriverSQLite {
t.Errorf("默认驱动应为 sqlite,实际 %s", c.Env.DBDriver)
}
if c.SQLitePath() != DefaultSQLitePath {
t.Errorf("SQLite 默认路径 = %s", c.SQLitePath())
}
if c.Env.Listen != ":8466" {
t.Errorf("默认监听地址错误: %s", c.Env.Listen)
}
if c.Env.StorageEngine != "local" {
t.Errorf("默认存储引擎错误: %s", c.Env.StorageEngine)
}
if c.Env.RedisAddr != "" {
t.Errorf("RedisAddr 应为空: %s", c.Env.RedisAddr)
}
}
func TestNewPostgresRequiresDSN(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "postgres")
t.Setenv("FCB_DB_DSN", "")
if _, err := New(); err == nil {
t.Fatal("postgres 模式缺少 FCB_DB_DSN 应报错")
}
t.Setenv("FCB_DB_DSN", "postgres://user:pass@localhost:5432/fcb")
c, err := New()
if err != nil {
t.Fatalf("postgres + DSN 应可构造: %v", err)
}
if c.Env.DBDriver != DBDriverPostgres {
t.Errorf("驱动应为 postgres,实际 %s", c.Env.DBDriver)
}
if c.SQLitePath() != "" {
t.Errorf("postgres 模式 SQLitePath 应为空: %s", c.SQLitePath())
}
}
func TestNewInvalidDriver(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "mysql")
t.Setenv("FCB_DB_DSN", "x")
if _, err := New(); err == nil {
t.Fatal("非法驱动应报错")
}
}
func TestNewInvalidEngine(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
t.Setenv("FCB_STORAGE_ENGINE", "onedrive")
if _, err := New(); err == nil {
t.Fatal("非法引擎应报错")
}
}
func TestEnvOverridesAndDefaults(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
t.Setenv("FCB_LISTEN", ":9999")
t.Setenv("FCB_STORAGE_ENGINE", "webdav")
c, err := New()
if err != nil {
t.Fatalf("New 失败: %v", err)
}
if c.Env.Listen != ":9999" || c.Env.StorageEngine != "webdav" {
t.Fatalf("env 覆盖失败: %+v", c.Env)
}
// 默认值对齐参考 DEFAULT_CONFIG
if got := c.GetInt("uploadSize"); got != 1024*1024*10 {
t.Errorf("uploadSize 默认值 = %d", got)
}
if got := c.GetInt("errorCount"); got != 10 {
t.Errorf("errorCount 默认值 = %d", got)
}
if got := c.GetInt("loginCount"); got != 5 {
t.Errorf("loginCount 默认值 = %d", got)
}
if got := c.GetInt("loginMinute"); got != 15 {
t.Errorf("loginMinute 默认值 = %d", got)
}
if got := c.GetBool("openUpload"); !got {
t.Error("openUpload 默认应为开启")
}
if c.EnableChunk() {
t.Error("enableChunk 默认应关闭")
}
// 新增字段(需求 ①)
if c.LogoURL() != DefaultLogoURL {
t.Errorf("logo_url 默认值 = %s", c.LogoURL())
}
if c.FaviconURL() != DefaultFaviconURL {
t.Errorf("favicon_url 默认值 = %s", c.FaviconURL())
}
if c.SiteName() == "" {
t.Error("site_name 默认值不应为空")
}
// 过期方式与文件类型
if len(c.ExpireStyle()) != 5 {
t.Errorf("expireStyle 默认值 = %v", c.ExpireStyle())
}
if len(c.AllowedFileTypes()) != 1 || c.AllowedFileTypes()[0] != "*" {
t.Errorf("allowed_file_types 默认值 = %v", c.AllowedFileTypes())
}
}
func TestKVOverridesEnvAndDefaults(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
c, _ := New()
c.ApplyKV(map[string]any{
"uploadSize": 1024,
"openUpload": 0,
"site_name": "我的快递柜",
"logo_url": "https://example.com/logo.svg",
"internalKey": "x", // 非下划线开头允许;下划线开头被拒
"_secret": "no",
})
if got := c.GetInt("uploadSize"); got != 1024 {
t.Errorf("KV 覆盖 uploadSize 失败: %d", got)
}
if c.OpenUpload() {
t.Error("KV 覆盖 openUpload 失败")
}
if c.SiteName() != "我的快递柜" {
t.Errorf("site_name KV 覆盖失败: %s", c.SiteName())
}
if c.LogoURL() != "https://example.com/logo.svg" {
t.Errorf("logo_url KV 覆盖失败: %s", c.LogoURL())
}
if _, ok := c.Get("_secret"); ok {
t.Error("下划线内部键不应可通过 ApplyKV 覆盖")
}
}
func TestAdminSessionExpireClamp(t *testing.T) {
t.Setenv("FCB_DB_DSN", "postgres://x")
c, _ := New()
if got := c.AdminSessionExpireSeconds(); got != AdminSessionExpireDefault {
t.Errorf("默认会话有效期 = %d", got)
}
c.ApplyKV(map[string]any{"adminSessionExpire": 7 * 24 * 60 * 60})
if got := c.AdminSessionExpireSeconds(); got != 7*24*60*60 {
t.Errorf("7 天会话有效期 = %d", got)
}
c.ApplyKV(map[string]any{"adminSessionExpire": 3600}) // 非整天,回落默认
if got := c.AdminSessionExpireSeconds(); got != AdminSessionExpireDefault {
t.Errorf("非法值应回落默认 = %d", got)
}
}
+97
View File
@@ -0,0 +1,97 @@
// Package config — schema.go 定义 v2 新增配置键(KVschema
// 键名常量、类型、默认值与取值边界。管理与 API 层(t2)按下表读写与校验,
// 文档(t4)按本表生成说明。键名除参考实现既有 camelCase 键外,
// v2 新增键统一 snake_case。
package config
// —— v2 新增/沿用键名常量(单一事实来源;settings 包会 re-export)——
// 命名规则:v2 新增键 snake_case;与参考实现对齐的既有键保持原拼写。
const (
// 需求 ①:背景图
KeyBackground = "background" // 参考实现既有键(v1 兼容保留)
KeyBackgroundURL = "background_url" // v2 新增:背景图 URL 或上传后的访问地址(空=默认主题)
// 需求 ②:页脚
KeyFooterText = "footer_text" // v2 新增:页脚自定义内容(纯文本或受控 HTML 片段)
KeyFooterBeian = "footer_beian" // v2 新增:备案号(如 京ICP备2024xxxxxx号-1
// 需求 ③:系统通知
KeyNotifyEnabled = "notify_enabled" // v2 新增:通知开关,1 开启 / 0 关闭
KeyNotifyTitle = "notify_title" // 既有键:通知标题
KeyNotifyContent = "notify_content" // 既有键:通知内容(允许 <a> 等受控 HTML
// 需求 ④:保存策略(上传页动态读取并在范围内选择)
KeyMaxSaveSeconds = "max_save_seconds" // 既有键:最长保存秒数,0=不限制(仅受默认 7 天兜底)
KeyMaxSaveCount = "max_save_count" // v2 新增:单次分享最大可取(保存)次数上限,0=不限制
KeyExpireStyle = "expireStyle" // 既有键:允许的过期方式白名单
// 需求 ④:上传频率限制(既有键,对齐参考 ip_limit["upload"]
KeyUploadCount = "uploadCount" // 窗口内允许上传次数
KeyUploadMinute = "uploadMinute" // 频率窗口(分钟)
// 需求 ④⑩:存储策略(最大文件大小/允许类型/总容量)
KeyUploadSize = "uploadSize" // 既有键:单文件上限(字节),参考实现语义
KeyMaxFileSize = "max_file_size" // v2 新增:存储策略-单文件上限(字节),0=回落 uploadSize
KeyAllowedTypes = "allowed_file_types" // 既有键:允许类型白名单("*" 不限制)
KeyStorageLimit = "storageLimit" // 既有键:站点总容量(字节),0=不限制
KeyOpenUpload = "openUpload" // 既有键:游客上传开关
// v3:存储引擎运行时可配(热切换;file_storage 为参考既有键保留兼容)
KeyStorageEngine = "storage_engine" // 当前存储引擎:local|s3|webdav
KeySiteDomain = "site_domain" // 站点对外域名(空=分享链接用当前地址)
)
// —— 取值边界(管理端保存与 API 校验用)——
const (
// 保存时间上限:最长 365 天,0 表示不限制。
MaxSaveSecondsMax = 365 * 24 * 60 * 60
// 保存次数上限:最长 100000 次,0 表示不限制。
MaxSaveCountMax = 100000
// 单文件大小上限:最长 10 GiB0 表示回落 uploadSize。
MaxFileSizeMax = 10 * 1024 * 1024 * 1024
// 背景图 URL 最大长度(含 data: 之外的普通 http(s) URL)。
BackgroundURLMaxLen = 2048
// 页脚自定义内容最大长度。
FooterTextMaxLen = 2000
// 备案号最大长度。
FooterBeianMaxLen = 128
// 通知标题/内容最大长度。
NotifyTitleMaxLen = 128
NotifyContentMaxLen = 2000
)
// KVSchemaEntry 配置键元数据:类型 / 默认值 / 说明,供管理端 UI 与文档生成。
type KVSchemaEntry struct {
Key string // KV 键名
Type string // string | int | int64 | bool | []string
Default any // 默认值(与 defaults() 保持一致,测试保证同步)
Min int64 // 数值键最小值(字符串键为长度下界)
Max int64 // 数值键最大值(字符串键为长度上界;-1 不限制)
Description string // 中文说明
}
// KVSchema v2 全量配置键 schema 表(含既有策略键,供管理端/文档/AI 校验)。
// 注意:Default 与 config defaults() 逐一对应(schema_test 保证)。
func KVSchema() []KVSchemaEntry {
return []KVSchemaEntry{
// —— 需求 ① 背景图 ——
{KeyBackgroundURL, "string", "", 0, BackgroundURLMaxLen, "背景图 URL 或上传后地址(空=主题默认)"},
// —— 需求 ② 页脚 ——
{KeyFooterText, "string", "", 0, FooterTextMaxLen, "页脚自定义内容(纯文本或受控 HTML 片段)"},
{KeyFooterBeian, "string", "", 0, FooterBeianMaxLen, "备案号,展示于页脚"},
// —— 需求 ③ 系统通知 ——
{KeyNotifyEnabled, "int", 1, 0, 1, "系统通知开关:1 右上角悬浮窗展示 / 0 关闭"},
{KeyNotifyTitle, "string", "系统通知", 0, NotifyTitleMaxLen, "通知标题"},
{KeyNotifyContent, "string", "欢迎使用文件快传,拖拽或粘贴即可分享文本与文件。", 0, NotifyContentMaxLen, "通知内容(允许 <a> 等受控 HTML"},
// —— 需求 ④ 保存策略 ——
{KeyMaxSaveSeconds, "int64", int64(0), 0, MaxSaveSecondsMax, "最长保存秒数上限,0=不限制(默认 7 天兜底)"},
{KeyMaxSaveCount, "int", 0, 0, MaxSaveCountMax, "单次分享最大可取次数上限,0=不限制"},
{KeyExpireStyle, "[]string", []string{"day", "hour", "minute", "forever", "count"}, -1, -1, "上传页可选过期方式白名单"},
// —— 需求 ④ 上传频率限制(既有键对齐参考)——
{KeyUploadCount, "int", 10, 1, 10000, "频率窗口内允许的上传次数"},
{KeyUploadMinute, "int", 1, 1, 1440, "上传频率窗口(分钟)"},
// —— 需求 ④⑩ 存储策略 ——
{KeyMaxFileSize, "int64", int64(0), 0, MaxFileSizeMax, "存储策略-单文件上限(字节),0=回落 uploadSize"},
{KeyUploadSize, "int64", int64(1024 * 1024 * 10), 1024, MaxFileSizeMax, "单文件上限(字节),参考实现语义"},
{KeyAllowedTypes, "[]string", []string{"*"}, -1, -1, "允许上传类型白名单(\"*\" 不限制)"},
{KeyStorageLimit, "int64", int64(0), 0, -1, "站点总容量(字节),0=不限制"},
{KeyOpenUpload, "int", 1, 0, 1, "游客上传开关:1 开 / 0 需管理员登录"},
// —— v3 存储引擎(热切换;引擎参数键沿用 defaults() 既有键,管理端经 config get/update 读写)——
{KeyStorageEngine, "string", "", 0, 16, "当前存储引擎:local|s3|webdav(热切换,健康检查通过才生效;空=回落启动值 FCB_STORAGE_ENGINE"},
{KeySiteDomain, "string", "", 0, 256, "站点对外域名(http(s)://host[:port],不带路径;空=分享链接用当前访问地址)"},
}
}
+91
View File
@@ -0,0 +1,91 @@
// schema 同步测试:保证 config.KVSchema() 的默认值/键集合与 defaults() 完全一致,
// 与 settings 包 re-export 的键名常量同源。新增键时任何一处漏改都会在此失败。
package config
import (
"encoding/json"
"testing"
)
// TestKVSchemaDefaultsMatchDefaults KVSchema 的 Default 必须 === defaults() 中同名键。
func TestKVSchemaDefaultsMatchDefaults(t *testing.T) {
def := defaults()
for _, e := range KVSchema() {
want, ok := def[e.Key]
if !ok {
t.Fatalf("schema 键 %q 缺少 defaults() 默认值", e.Key)
}
// 类型规范化比较(JSON 序列化可比较 []string / int / float
a, _ := json.Marshal(e.Default)
b, _ := json.Marshal(want)
if string(a) != string(b) {
t.Fatalf("键 %q 默认值不一致: schema=%s defaults=%s", e.Key, a, b)
}
}
}
// TestKVSchemaNoDuplicates 键名不得重复。
func TestKVSchemaNoDuplicates(t *testing.T) {
seen := map[string]bool{}
for _, e := range KVSchema() {
if seen[e.Key] {
t.Fatalf("schema 键 %q 重复定义", e.Key)
}
seen[e.Key] = true
}
}
// TestV2NewKeysPresent v2 新增键必须在 schema 与 defaults 中同时存在。
func TestV2NewKeysPresent(t *testing.T) {
def := defaults()
newKeys := []string{
KeyBackgroundURL, KeyFooterText, KeyFooterBeian,
KeyNotifyEnabled, KeyMaxSaveCount, KeyMaxFileSize,
}
for _, k := range newKeys {
if _, ok := def[k]; !ok {
t.Fatalf("v2 新键 %q 缺少默认值", k)
}
}
}
// TestV2AccessorDefaults v2 便捷访问器默认语义。
func TestV2AccessorDefaults(t *testing.T) {
t.Setenv("FCB_DB_DRIVER", "sqlite")
t.Setenv("FCB_DB_DSN", "")
c, err := New()
if err != nil {
t.Fatalf("New: %v", err)
}
// 背景图:background_url 与 background 均空 → 空
if c.BackgroundURL() != "" {
t.Fatalf("背景图默认应为空: %q", c.BackgroundURL())
}
// legacy background 键兜底
c.ApplyKV(map[string]any{KeyBackground: "/legacy/bg.jpg"})
if c.BackgroundURL() != "/legacy/bg.jpg" {
t.Fatalf("legacy background 应回落生效: %q", c.BackgroundURL())
}
// max_file_size > 0 时优先于 uploadSize
c.ApplyKV(map[string]any{KeyMaxFileSize: int64(1024), KeyUploadSize: int64(2048)})
if c.MaxFileSize() != 1024 {
t.Fatalf("max_file_size 应优先: %d", c.MaxFileSize())
}
// max_file_size = 0 回落 uploadSize
c.ApplyKV(map[string]any{KeyMaxFileSize: int64(0)})
if c.MaxFileSize() != 2048 {
t.Fatalf("max_file_size=0 应回落 uploadSize: %d", c.MaxFileSize())
}
// 通知默认开启
if !c.NotifyEnabled() {
t.Fatal("notify_enabled 默认应开启")
}
// 保存次数上限默认不限制
if c.MaxSaveCount() != 0 {
t.Fatalf("max_save_count 默认应 0: %d", c.MaxSaveCount())
}
// 页脚默认空
if c.FooterText() != "" || c.FooterBeian() != "" {
t.Fatal("页脚默认应为空")
}
}
+139
View File
@@ -0,0 +1,139 @@
// Package database 负责数据库连接与迁移(需求 ⑧:双方言):
// - sqlite(默认):modernc.org/sqlite 纯 Go 驱动(GORM 封装 glebarez/sqlite),零 CGO、零外部依赖;
// - postgres:可选,配置 FCB_DB_DRIVER=postgres + FCB_DB_DSN 后启用。
//
// 两方言共用 GORM 抽象层,AutoMigrate 与全部业务查询保持方言无关;
// 唯一的原生 SQL(migrates 建表)已改为双方言分支。
package database
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"time"
"github.com/glebarez/sqlite"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"filecodebox/internal/config"
"filecodebox/internal/model"
)
// Options 连接选项(main.go 从 config.Env 装配)。
type Options struct {
Driver string // sqlite | postgres(空按 sqlite 处理)
DSN string // postgres 连接串;sqlite 为文件路径(空回退 config.DefaultSQLitePath
}
// Open 按驱动连接数据库并执行连接池设置与探活。
func Open(ctx context.Context, opts Options) (*gorm.DB, error) {
driver := strings.ToLower(strings.TrimSpace(opts.Driver))
if driver == "" {
driver = config.DBDriverSQLite
}
var dialector gorm.Dialector
switch driver {
case config.DBDriverSQLite:
path := strings.TrimSpace(opts.DSN)
if path == "" {
path = config.DefaultSQLitePath
}
// 自动创建父目录(如 ./data),对齐参考实现 data_root 语义
if dir := filepath.Dir(path); dir != "" && dir != "." {
if err := os.MkdirAll(dir, 0o755); err != nil {
return nil, fmt.Errorf("database: 创建 SQLite 目录 %s 失败: %w", dir, err)
}
}
// DSN 参数:busy_timeout 防写锁竞态;WAL 提升并发读写(Query 参数形式,驱动原生支持)
dsn := path + "?_pragma=busy_timeout(10000)&_pragma=journal_mode(WAL)&_pragma=foreign_keys(1)"
dialector = sqlite.Open(dsn)
case config.DBDriverPostgres:
if strings.TrimSpace(opts.DSN) == "" {
return nil, fmt.Errorf("database: FCB_DB_DRIVER=postgres 需要提供 FCB_DB_DSN")
}
dialector = postgres.Open(opts.DSN)
default:
return nil, fmt.Errorf("database: 不支持的数据库驱动 %q(仅支持 sqlite|postgres", driver)
}
db, err := gorm.Open(dialector, &gorm.Config{
Logger: logger.Default.LogMode(logger.Warn),
// 避免 GORM 生成方言特有子句;时间语义由应用层统一(容器本地时区)
NowFunc: time.Now,
})
if err != nil {
return nil, fmt.Errorf("database: 连接 %s 失败: %w", driver, err)
}
sqlDB, err := db.DB()
if err != nil {
return nil, err
}
// 连接池:SQLite 单文件场景保守设置;Postgres 沿用 v1 参数
switch driver {
case config.DBDriverSQLite:
sqlDB.SetMaxOpenConns(8)
sqlDB.SetMaxIdleConns(4)
sqlDB.SetConnMaxLifetime(0) // 长连接文件句柄,无需轮换
case config.DBDriverPostgres:
sqlDB.SetMaxOpenConns(32)
sqlDB.SetMaxIdleConns(8)
sqlDB.SetConnMaxLifetime(time.Hour)
}
// 连接探活(带超时)
pingCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
if err := sqlDB.PingContext(pingCtx); err != nil {
return nil, fmt.Errorf("database: %s 探活失败: %w", driver, err)
}
return db, nil
}
// Migrate 执行迁移:先建迁移台账表(双方言分支),再 AutoMigrate 全部模型。
func Migrate(ctx context.Context, db *gorm.DB) error {
if err := createMigratesTable(ctx, db); err != nil {
return err
}
if err := model.AutoMigrate(db); err != nil {
return fmt.Errorf("database: AutoMigrate 失败: %w", err)
}
return nil
}
// createMigratesTable 创建迁移台账表。
// 双方言差异:自增主键 postgres 用 BIGSERIAL、sqlite 用 INTEGER PRIMARY KEY AUTOINCREMENT
// 时间戳默认值 postgres 用 CURRENT_TIMESTAMP、sqlite 用 CURRENT_TIMESTAMP(等价)。
func createMigratesTable(ctx context.Context, db *gorm.DB) error {
ddl := `
CREATE TABLE IF NOT EXISTS migrates (
id INTEGER PRIMARY KEY AUTOINCREMENT,
migration_file VARCHAR(255) NOT NULL UNIQUE,
executed_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)`
if db.Dialector.Name() == config.DBDriverPostgres {
ddl = `
CREATE TABLE IF NOT EXISTS migrates (
id BIGSERIAL PRIMARY KEY,
migration_file VARCHAR(255) NOT NULL UNIQUE,
executed_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)`
}
if err := db.WithContext(ctx).Exec(ddl).Error; err != nil {
return fmt.Errorf("database: 创建 migrates 表失败: %w", err)
}
return nil
}
// Close 关闭底层连接。
func Close(db *gorm.DB) error {
sqlDB, err := db.DB()
if err != nil {
return err
}
return sqlDB.Close()
}
+237
View File
@@ -0,0 +1,237 @@
// 数据库双方言测试(需求 ⑧):
// - sqlite:始终执行(纯 Go,临时目录建库);
// - postgres:设置 FCB_TEST_PG_DSN(真实连接串)后执行,未设置时跳过。
//
// 覆盖:Open/Migrate 全表建立、settings KV 读写、JSON 字段往返、
// 分页查询(LIMIT/OFFSET 语义)、布尔/时间字段往返 —— 双方言逐项比对。
package database_test
import (
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"gorm.io/gorm"
"filecodebox/internal/database"
"filecodebox/internal/model"
)
// pgTestDSN 返回 Postgres 测试连接串;未设置 FCB_TEST_PG_DSN 时返回空。
func pgTestDSN(t *testing.T) string {
t.Helper()
dsn := os.Getenv("FCB_TEST_PG_DSN")
if dsn == "" {
t.Skip("未设置 FCB_TEST_PG_DSN,跳过 Postgres 双方言用例(sqlite 用例仍执行)")
}
return dsn
}
// openTestDB 按方言打开数据库并执行迁移;返回 gorm 实例与关闭函数。
func openTestDB(t *testing.T, driver, dsn string) (*gorm.DB, func()) {
t.Helper()
if dsn == "" {
// sqlite:临时文件库
dir := t.TempDir()
dsn = filepath.Join(dir, "test.db")
}
ctx := context.Background()
db, err := database.Open(ctx, database.Options{Driver: driver, DSN: dsn})
if err != nil {
t.Fatalf("[%s] Open 失败: %v", driver, err)
}
if err := database.Migrate(ctx, db); err != nil {
_ = database.Close(db)
t.Fatalf("[%s] Migrate 失败: %v", driver, err)
}
return db, func() { _ = database.Close(db) }
}
// runDialectSuite 双方言共用的行为断言集。
func runDialectSuite(t *testing.T, db *gorm.DB) {
t.Helper()
ctx := context.Background()
// —— 1. 全表建立 ——
for _, m := range model.AllModels() {
if !db.Migrator().HasTable(m) {
t.Fatalf("表 %T 未创建", m)
}
}
// —— 2. settings KV 读写 + JSON 字段往返 ——
// GORM 软特性:KeyValue.Value 为 *stringJSON 文本),双方言 text 类型
// 可重跑:先清掉同键旧行(共享测试库场景)
if err := db.WithContext(ctx).Where(model.KeyValue{Key: "settings"}).Delete(&model.KeyValue{}).Error; err != nil {
t.Fatalf("KV 旧数据清理失败: %v", err)
}
kv := map[string]any{"background_url": "https://example.com/bg.jpg", "footer_beian": "京ICP备2024000001号-1", "max_save_seconds": 3600}
raw, err := json.Marshal(kv)
if err != nil {
t.Fatalf("marshal KV: %v", err)
}
row := model.KeyValue{Key: "settings", Value: strPtr(string(raw))}
if err := db.WithContext(ctx).Create(&row).Error; err != nil {
t.Fatalf("KV 写入失败: %v", err)
}
var got model.KeyValue
if err := db.WithContext(ctx).Where(model.KeyValue{Key: "settings"}).First(&got).Error; err != nil {
t.Fatalf("KV 读取失败: %v", err)
}
parsed := map[string]any{}
if err := json.Unmarshal([]byte(*got.Value), &parsed); err != nil {
t.Fatalf("KV JSON 解析失败: %v", err)
}
if parsed["background_url"] != "https://example.com/bg.jpg" {
t.Fatalf("KV JSON 字段往返不一致: %v", parsed)
}
// 更新(先查后改,方言无关)
if err := db.WithContext(ctx).Model(&got).Update("value", strPtr(`{"notify_enabled":0}`)).Error; err != nil {
t.Fatalf("KV 更新失败: %v", err)
}
var got2 model.KeyValue
_ = db.WithContext(ctx).Where(model.KeyValue{Key: "settings"}).First(&got2)
if *got2.Value != `{"notify_enabled":0}` {
t.Fatalf("KV 更新未生效: %s", *got2.Value)
}
// —— 3. 分页查询(LIMIT/OFFSET)——
// 每次运行用随机前缀避免脏数据互相影响
prefix := fmt.Sprintf("pg%d_", time.Now().UnixNano())
for i := 0; i < 25; i++ {
fc := model.FileCodes{
Code: fmt.Sprintf("%s%03d", prefix, i),
ExpiredCount: -1,
IsChunked: i%2 == 0, // 布尔字段往返
}
if err := db.WithContext(ctx).Create(&fc).Error; err != nil {
t.Fatalf("FileCodes 写入失败: %v", err)
}
}
var page []model.FileCodes
if err := db.WithContext(ctx).
Where("code LIKE ?", prefix+"%").
Order("id ASC").
Limit(10).Offset(20).
Find(&page).Error; err != nil {
t.Fatalf("分页查询失败: %v", err)
}
if len(page) != 5 {
t.Fatalf("第二页应剩 5 条,实际 %d", len(page))
}
if page[0].Code != prefix+"020" {
t.Fatalf("分页偏移错误: %s", page[0].Code)
}
var total int64
if err := db.WithContext(ctx).Model(&model.FileCodes{}).
Where("code LIKE ?", prefix+"%").Count(&total).Error; err != nil {
t.Fatalf("计数查询失败: %v", err)
}
if total != 25 {
t.Fatalf("总数应 25,实际 %d", total)
}
// —— 4. 布尔/时间/可空字段往返 ——
now := time.Now().Truncate(time.Second) // sqlite 秒级精度
fc := model.FileCodes{
Code: prefix + "special",
ExpiredAt: &now,
ExpiredCount: 5,
Text: strPtr("你好 FileCodeBox"),
FileHash: strPtr("abc123"),
IsChunked: true,
}
if err := db.WithContext(ctx).Create(&fc).Error; err != nil {
t.Fatalf("完整字段写入失败: %v", err)
}
var back model.FileCodes
if err := db.WithContext(ctx).Where(model.FileCodes{Code: fc.Code}).First(&back).Error; err != nil {
t.Fatalf("完整字段读取失败: %v", err)
}
if back.Text == nil || *back.Text != "你好 FileCodeBox" {
t.Fatalf("text 字段往返不一致: %v", back.Text)
}
if !back.IsChunked {
t.Fatal("布尔字段往返不一致")
}
if back.ExpiredAt == nil {
t.Fatal("时间字段往返丢失")
}
if diff := back.ExpiredAt.Sub(now); diff > time.Second || diff < -time.Second {
t.Fatalf("时间字段偏差过大: %v", diff)
}
if back.FileHash == nil || *back.FileHash != "abc123" {
t.Fatalf("可空字段往返不一致: %v", back.FileHash)
}
// LOWER + LIKEadmin 列表检索路径:真实代码先对关键词小写化再拼 LIKE 模式,
// 对齐 admin.go 的 "LOWER(code) LIKE ?" 用法,双方言均支持)
var hits int64
lowerPattern := "%" + strings.ToLower(prefix+"SPECIAL") + "%"
if err := db.WithContext(ctx).Model(&model.FileCodes{}).
Where("LOWER(code) LIKE ?", lowerPattern).Count(&hits).Error; err != nil {
t.Fatalf("LOWER/LIKE 查询失败: %v", err)
}
if hits != 1 {
t.Fatalf("LOWER/LIKE 命中数应 1,实际 %d", hits)
}
// 可重跑:清理本前缀数据(共享测试库场景)
if err := db.WithContext(ctx).Where("code LIKE ?", prefix+"%").Delete(&model.FileCodes{}).Error; err != nil {
t.Fatalf("清理测试数据失败: %v", err)
}
}
func strPtr(s string) *string { return &s }
// TestSQLiteDialect sqlite(默认模式):临时文件库全流程。
func TestSQLiteDialect(t *testing.T) {
db, closeFn := openTestDB(t, "sqlite", "")
defer closeFn()
runDialectSuite(t, db)
}
// TestSQLiteInMemoryDialect sqlite 内存库(DSN 为 :memory: 等价路径场景)。
func TestSQLiteInMemoryDialect(t *testing.T) {
dir := t.TempDir()
db, closeFn := openTestDB(t, "sqlite", filepath.Join(dir, "mem.db"))
defer closeFn()
runDialectSuite(t, db)
}
// TestPostgresDialect postgres(可选模式):FCB_TEST_PG_DSN 指向真实实例。
func TestPostgresDialect(t *testing.T) {
dsn := pgTestDSN(t)
db, closeFn := openTestDB(t, "postgres", dsn)
defer closeFn()
runDialectSuite(t, db)
}
// TestOpenRejectsUnknownDriver 非法驱动应报错。
func TestOpenRejectsUnknownDriver(t *testing.T) {
if _, err := database.Open(context.Background(), database.Options{Driver: "mysql", DSN: "x"}); err == nil {
t.Fatal("非法驱动应报错")
}
}
// TestOpenPostgresRequiresDSN postgres 模式缺 DSN 应报错。
func TestOpenPostgresRequiresDSN(t *testing.T) {
if _, err := database.Open(context.Background(), database.Options{Driver: "postgres", DSN: ""}); err == nil {
t.Fatal("postgres 缺 DSN 应报错")
}
}
// TestSQLiteAutoCreatesDataDir sqlite 默认相对路径下自动创建父目录。
func TestSQLiteAutoCreatesDataDir(t *testing.T) {
dir := t.TempDir()
nested := filepath.Join(dir, "deep", "data", "fcb.db")
db, closeFn := openTestDB(t, "sqlite", nested)
defer closeFn()
if _, err := os.Stat(nested); err != nil {
t.Fatalf("数据库文件应已创建: %v", err)
}
runDialectSuite(t, db)
}
+124
View File
@@ -0,0 +1,124 @@
// Package janitor 后台清理循环(安全审计 M5):
// 回收过期容量预留、超时未完成的上传会话(含其分片对象)与过期预签名会话
// (direct 模式残留对象一并删除)。此前这些资源仅在同 token 复用/显式取消时
// 释放,恶意 init 可长期占用容量预留或累积垃圾数据。
package janitor
import (
"context"
"errors"
"log"
"time"
"gorm.io/gorm"
"filecodebox/internal/model"
"filecodebox/internal/storage"
)
// chunkSessionMaxAge 未完成分片会话的最大保留时长(预留 TTL 为 2h,
// 会话保留 24h 以支持断点续传;超时后由本循环清理)。
const chunkSessionMaxAge = 24 * time.Hour
// presignGrace 过期预签名会话的宽限时长(到点即删,避免与在途 confirm 竞争)。
const presignGrace = time.Hour
// Start 启动周期清理循环;ctx 取消时退出。
func Start(ctx context.Context, db *gorm.DB, store *storage.Manager, interval time.Duration) {
go func() {
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
Run(ctx, db, store)
}
}
}()
}
// Run 执行一轮清理;单项失败仅记日志,不影响其他项。
func Run(ctx context.Context, db *gorm.DB, store *storage.Manager) {
now := time.Now()
cleanExpiredReservations(ctx, db, now)
cleanExpiredChunkSessions(ctx, db, store, now)
cleanExpiredPresignSessions(ctx, db, store, now)
}
// cleanExpiredReservations 删除全部过期容量预留。
func cleanExpiredReservations(ctx context.Context, db *gorm.DB, now time.Time) {
if err := db.WithContext(ctx).
Where("expires_at <= ?", now).
Delete(&model.StorageReservation{}).Error; err != nil {
log.Printf("[janitor] 清理过期容量预留失败: %v", err)
}
}
// engineFor 按归属引擎取回实例;空/未知引擎回落当前引擎(对齐 API 层 storeFor 语义)。
func engineFor(store *storage.Manager, name string) (storage.Storage, error) {
if name != "" && storage.ValidEngine(name) {
if s, err := store.EngineOf(name); err == nil {
return s, nil
}
}
return store.Current(), nil
}
// cleanExpiredChunkSessions 清理超时未完成的分片会话及其分片对象。
func cleanExpiredChunkSessions(ctx context.Context, db *gorm.DB, store *storage.Manager, now time.Time) {
var sessions []model.UploadChunk
if err := db.WithContext(ctx).
Where("chunk_index = -1 AND created_at < ?", now.Add(-chunkSessionMaxAge)).
Limit(200).
Find(&sessions).Error; err != nil {
log.Printf("[janitor] 查询过期分片会话失败: %v", err)
return
}
for _, s := range sessions {
engine, err := engineFor(store, s.Engine)
if err == nil && s.SavePath != "" {
if err := engine.CleanChunks(ctx, s.UploadID, s.SavePath); err != nil &&
!errors.Is(err, storage.ErrNotFound) && !errors.Is(err, storage.ErrInvalidPath) {
log.Printf("[janitor] 清理分片对象失败 upload_id=%s: %v", s.UploadID, err)
}
}
if err := db.WithContext(ctx).
Where("upload_id = ?", s.UploadID).
Delete(&model.UploadChunk{}).Error; err != nil {
log.Printf("[janitor] 删除过期分片会话失败 upload_id=%s: %v", s.UploadID, err)
continue
}
log.Printf("[janitor] 已清理超时分片会话 upload_id=%s file=%s", s.UploadID, s.FileName)
}
}
// cleanExpiredPresignSessions 清理过期预签名会话;direct 模式残留对象一并删除。
func cleanExpiredPresignSessions(ctx context.Context, db *gorm.DB, store *storage.Manager, now time.Time) {
var sessions []model.PresignUploadSession
if err := db.WithContext(ctx).
Where("expires_at < ?", now.Add(-presignGrace)).
Limit(200).
Find(&sessions).Error; err != nil {
log.Printf("[janitor] 查询过期预签名会话失败: %v", err)
return
}
for _, s := range sessions {
if s.Mode == "direct" && s.SavePath != "" {
if engine, err := engineFor(store, s.Engine); err == nil {
if err := engine.DeleteFile(ctx, s.SavePath); err != nil &&
!errors.Is(err, storage.ErrNotFound) && !errors.Is(err, storage.ErrInvalidPath) {
log.Printf("[janitor] 删除直传残留对象失败 upload_id=%s: %v", s.UploadID, err)
}
}
}
if err := db.WithContext(ctx).
Where("upload_id = ?", s.UploadID).
Delete(&model.PresignUploadSession{}).Error; err != nil {
log.Printf("[janitor] 删除过期预签名会话失败 upload_id=%s: %v", s.UploadID, err)
continue
}
log.Printf("[janitor] 已清理过期预签名会话 upload_id=%s mode=%s", s.UploadID, s.Mode)
}
}
+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)
}
}
+152
View File
@@ -0,0 +1,152 @@
// Package model 定义 GORM 数据模型与 Postgres 自动迁移。
// 字段对齐参考实现 apps/base/models.py,并新增审计日志表。
package model
import (
"time"
"gorm.io/gorm"
)
// FileCodes 文件/文本分享记录(对齐参考 FileCodes)。
type FileCodes struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Code string `gorm:"column:code;size:255;uniqueIndex;not null" json:"code"` // 取件码
Prefix string `gorm:"size:255;default:''" json:"prefix"` // 文件名前缀/文本分享标记
Suffix string `gorm:"size:255;default:''" json:"suffix"` // 文件名后缀(含扩展名)
UUIDFileName *string `gorm:"size:255" json:"uuid_file_name"` // 存储侧 UUID 文件名
FilePath *string `gorm:"size:255" json:"file_path"` // 存储侧相对路径
Size int64 `gorm:"default:0" json:"size"` // 字节数;文本为字符数
Text *string `gorm:"type:text" json:"text"` // 文本分享内容
ExpiredAt *time.Time `json:"expired_at"` // 过期时间;永久分享为 NULL
ExpiredCount int `gorm:"default:0" json:"expired_count"` // 剩余可取次数;<0 表示按时间过期
UsedCount int `gorm:"default:0" json:"used_count"` // 已取次数
CreatedAt time.Time `json:"created_at"`
FileHash *string `gorm:"size:64" json:"file_hash"` // SHA256
IsChunked bool `gorm:"default:false" json:"is_chunked"`
UploadID *string `gorm:"size:36" json:"upload_id"` // 分片上传会话 ID
Engine string `gorm:"size:16;default:''" json:"engine"` // 归属存储引擎(v3local|s3|webdav;空=历史数据按当前引擎取)
}
// TableName 表名。
func (FileCodes) TableName() string { return "file_codes" }
// Expired 判断是否已过期(对齐参考语义:expired_count<0 按时间,否则按次数)。
func (f *FileCodes) Expired(now time.Time) bool {
if f.ExpiredAt == nil {
return false
}
if f.ExpiredCount < 0 {
return f.ExpiredAt.Before(now)
}
return f.ExpiredCount <= 0
}
// UploadChunk 分片上传记录(对齐参考 UploadChunk)。
type UploadChunk struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
UploadID string `gorm:"size:36;index:idx_upload_chunk,unique,priority:1;not null" json:"upload_id"`
ChunkIndex int `gorm:"index:idx_upload_chunk,unique,priority:2;not null" json:"chunk_index"`
ChunkHash string `gorm:"size:64;not null" json:"chunk_hash"` // 分片 SHA256
TotalChunks int `json:"total_chunks"`
FileSize int64 `json:"file_size"`
ChunkSize int `json:"chunk_size"`
FileName string `gorm:"size:255" json:"file_name"`
SavePath string `gorm:"size:512" json:"save_path"`
CreatedAt time.Time `json:"created_at"`
Completed bool `gorm:"default:false" json:"completed"`
Engine string `gorm:"size:16;default:''" json:"engine"` // 归属存储引擎(v3:分片与会话记录当时引擎,合并走同一引擎)
}
// TableName 表名。
func (UploadChunk) TableName() string { return "upload_chunks" }
// KeyValue 运行时配置键值(对齐参考 KeyValue)。value 存 JSON。
type KeyValue struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Key string `gorm:"size:255;uniqueIndex;not null" json:"key"`
Value *string `gorm:"type:text" json:"value"` // JSON 字符串
CreatedAt time.Time `json:"created_at"`
}
// TableName 表名。
func (KeyValue) TableName() string { return "key_values" }
// PresignUploadSession 预签名直传会话(对齐参考 PresignUploadSession)。
type PresignUploadSession struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
UploadID string `gorm:"size:36;uniqueIndex;not null" json:"upload_id"`
FileName string `gorm:"size:255" json:"file_name"`
FileSize int64 `json:"file_size"`
SavePath string `gorm:"size:512" json:"save_path"`
Mode string `gorm:"size:10" json:"mode"` // direct=客户端直传 | proxy=服务器代理
ExpireValue int `json:"expire_value"`
ExpireStyle string `gorm:"size:20;default:day" json:"expire_style"`
CreatedAt time.Time `json:"created_at"`
ExpiresAt time.Time `json:"expires_at"`
Engine string `gorm:"size:16;default:''" json:"engine"` // 归属存储引擎(v3:直传/代理完成走同一引擎取回)
}
// TableName 表名。
func (PresignUploadSession) TableName() string { return "presign_upload_sessions" }
// IsExpired 会话是否已过期。
func (p *PresignUploadSession) IsExpired(now time.Time) bool { return p.ExpiresAt.Before(now) }
// StorageReservation 上传容量预留(尚未写入 file_codes 的占位)。
type StorageReservation struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Token string `gorm:"size:64;uniqueIndex;not null" json:"token"`
Size int64 `json:"size"`
ExpiresAt time.Time `gorm:"index" json:"expires_at"`
}
// TableName 表名。
func (StorageReservation) TableName() string { return "storage_reservations" }
// 审计结果常量。
const (
AuditResultSuccess = "success" // 操作成功
AuditResultDenied = "denied" // 被拒绝(限流/鉴权/策略)
AuditResultFailed = "failed" // 执行失败(服务端/客户端错误)
)
// AuditLog 上传/下载审计日志(需求 ③)。
type AuditLog struct {
ID int64 `gorm:"primaryKey;autoIncrement" json:"id"`
Action string `gorm:"size:32;index" json:"action"` // upload | download
FileCode string `gorm:"size:64;index" json:"file_code"` // 取件码(上传时为生成的码)
FileName string `gorm:"size:255" json:"file_name"` // 原始文件名/文本标记
SizeBytes int64 `json:"size_bytes"` // 文件总字节数
TransferredBytes int64 `json:"transferred_bytes"` // 本次实际传输字节数
IP string `gorm:"size:64;index" json:"ip"`
UserAgent string `gorm:"size:512" json:"user_agent"`
DeviceOS string `gorm:"size:64" json:"device_os"` // Windows/macOS/Android/iOS/Linux/Unknown
DeviceBrowser string `gorm:"size:64" json:"device_browser"` // Chrome/Firefox/Safari/Edge/...
DeviceType string `gorm:"size:32" json:"device_type"` // desktop/mobile/tablet/bot/other
Actor string `gorm:"size:64" json:"actor"` // admin | guest
Result string `gorm:"size:16;index" json:"result"` // success | denied | failed
ErrorMsg string `gorm:"size:512" json:"error_msg"`
DurationMs int64 `json:"duration_ms"`
CreatedAt time.Time `gorm:"index" json:"created_at"` // 操作时间
}
// TableName 表名。
func (AuditLog) TableName() string { return "audit_logs" }
// AllModels 全部需要迁移的模型。
func AllModels() []any {
return []any{
&FileCodes{},
&UploadChunk{},
&KeyValue{},
&PresignUploadSession{},
&StorageReservation{},
&AuditLog{},
}
}
// AutoMigrate 在 Postgres 上建表/补列;服务启动时调用。
func AutoMigrate(db *gorm.DB) error {
return db.AutoMigrate(AllModels()...)
}
+25
View File
@@ -0,0 +1,25 @@
// Package response 提供统一响应封装:{"code":200,"msg":"...","data":...}。
package response
import (
"net/http"
"github.com/gin-gonic/gin"
)
// Body 统一响应体。
type Body struct {
Code int `json:"code"`
Msg string `json:"msg"`
Data any `json:"data,omitempty"`
}
// OK 成功响应(code=200)。
func OK(c *gin.Context, data any) {
c.JSON(http.StatusOK, Body{Code: 200, Msg: "ok", Data: data})
}
// Fail 失败响应,httpStatus 与 code 语义一致(404 过期/不存在、403 拒绝、429 限流、500 服务端错误)。
func Fail(c *gin.Context, httpStatus int, msg string) {
c.AbortWithStatusJSON(httpStatus, Body{Code: httpStatus, Msg: msg})
}
+150
View File
@@ -0,0 +1,150 @@
// settings 包双方言测试:Manager 全流程(ensure 行、KV 读写合并、Reload、
// UpdateKV 屏蔽内部键、SystemStart)分别在 sqlite(默认)与 postgresFCB_TEST_PG_DSN)上执行。
package settings_test
import (
"context"
"os"
"path/filepath"
"testing"
"gorm.io/gorm"
"filecodebox/internal/config"
"filecodebox/internal/database"
"filecodebox/internal/settings"
)
// pgTestDSN 返回 Postgres 测试连接串;未设置 FCB_TEST_PG_DSN 时跳过。
func pgTestDSN(t *testing.T) string {
t.Helper()
dsn := os.Getenv("FCB_TEST_PG_DSN")
if dsn == "" {
t.Skip("未设置 FCB_TEST_PG_DSN,跳过 Postgres 双方言用例")
}
return dsn
}
// newTestManager 按方言构造库 + Manager(已完成 Migrate)。
func newTestManager(t *testing.T, driver, dsn string) (*settings.Manager, *gorm.DB, func()) {
t.Helper()
if dsn == "" {
dsn = filepath.Join(t.TempDir(), "settings-test.db")
}
ctx := context.Background()
db, err := database.Open(ctx, database.Options{Driver: driver, DSN: dsn})
if err != nil {
t.Fatalf("[%s] Open: %v", driver, err)
}
if err := database.Migrate(ctx, db); err != nil {
_ = database.Close(db)
t.Fatalf("[%s] Migrate: %v", driver, err)
}
// Config 仅作内存载体(驱动不回连数据库),postgres 模式给占位 DSN 以通过校验
t.Setenv("FCB_DB_DRIVER", driver)
if driver == "postgres" {
t.Setenv("FCB_DB_DSN", dsn)
} else {
t.Setenv("FCB_DB_DSN", "")
}
cfg, err := config.New()
if err != nil {
_ = database.Close(db)
t.Fatalf("[%s] config.New: %v", driver, err)
}
mgr, err := settings.NewManager(ctx, db, cfg)
if err != nil {
_ = database.Close(db)
t.Fatalf("[%s] NewManager: %v", driver, err)
}
return mgr, db, func() { _ = database.Close(db) }
}
// runManagerSuite 双方言共用的 Manager 行为断言。
func runManagerSuite(t *testing.T, mgr *settings.Manager) {
t.Helper()
ctx := context.Background()
// 1. 初始:未初始化(admin_token 空)
if mgr.IsInitialized() {
t.Fatal("初始 admin_token 为空应视为未初始化")
}
// 2. UpdateKV 写入策略键 → Reload 后读取生效
patch := map[string]any{
settings.KeyBackgroundURL: "https://example.com/bg.png",
settings.KeyFooterText: "自建部署,仅供内部演示",
settings.KeyFooterBeian: "京ICP备2024000001号-1",
settings.KeyNotifyEnabled: 0,
settings.KeyMaxSaveSeconds: 86400,
"_internal_secret": "must-drop", // 下划线内部键必须被拒
}
if err := mgr.UpdateKV(ctx, patch); err != nil {
t.Fatalf("UpdateKV 失败: %v", err)
}
if err := mgr.Reload(ctx); err != nil {
t.Fatalf("Reload 失败: %v", err)
}
cfg := mgr.Get()
if got := cfg.GetString(settings.KeyBackgroundURL); got != "https://example.com/bg.png" {
t.Fatalf("background_url 未生效: %q", got)
}
if got := cfg.GetString(settings.KeyFooterBeian); got != "京ICP备2024000001号-1" {
t.Fatalf("footer_beian 未生效: %q", got)
}
if cfg.GetBool(settings.KeyNotifyEnabled) {
t.Fatal("notify_enabled=0 应生效")
}
if got := cfg.MaxSaveSeconds(); got != 86400 {
t.Fatalf("max_save_seconds 未生效: %d", got)
}
if _, ok := cfg.Get("_internal_secret"); ok {
t.Fatal("下划线内部键不应进入运行时配置")
}
// 3. KV 合并语义:二次 UpdateKV 不覆盖未提及键
if err := mgr.UpdateKV(ctx, map[string]any{settings.KeyNotifyEnabled: 1}); err != nil {
t.Fatalf("二次 UpdateKV: %v", err)
}
if err := mgr.Reload(ctx); err != nil {
t.Fatalf("二次 Reload: %v", err)
}
cfg = mgr.Get()
if !cfg.GetBool(settings.KeyNotifyEnabled) {
t.Fatal("notify_enabled 二次写入应生效")
}
if got := cfg.GetString(settings.KeyFooterText); got == "" {
t.Fatal("二次写入不应清空 footer_text")
}
// 4. SystemStartsys_start 键写入且为毫秒时间戳
mgr.SystemStart(ctx)
// 5. 敏感键判定(双模式一致)
if !settings.IsSensitiveKey("admin_token") || !settings.IsSensitiveKey("jwt_secret") {
t.Fatal("admin_token/jwt_secret 应为敏感键")
}
if settings.IsSensitiveKey("footer_text") {
t.Fatal("footer_text 不应为敏感键")
}
// 6. KV schema 表完整性:全部键可从默认值读取
for _, e := range settings.KVSchema() {
if _, ok := cfg.Get(e.Key); !ok {
t.Fatalf("schema 键 %q 在默认配置中不存在", e.Key)
}
}
}
func TestManagerSQLite(t *testing.T) {
mgr, _, closeFn := newTestManager(t, "sqlite", "")
defer closeFn()
runManagerSuite(t, mgr)
}
func TestManagerPostgres(t *testing.T) {
dsn := pgTestDSN(t)
mgr, _, closeFn := newTestManager(t, "postgres", dsn)
defer closeFn()
runManagerSuite(t, mgr)
}
+98
View File
@@ -0,0 +1,98 @@
// Package settings 密码哈希与校验:
// 新密码使用 bcrypt(格式 bcrypt$<bcrypt原生哈希串>);同时兼容两代旧格式——
// sha256$salt$hash(上一版)与旧版明文(迁移校验)。
// 安全审计 M1:单轮 SHA256+盐抗 GPU 爆破不足,新哈希统一升级 bcrypt。
// 兼容策略:VerifyPassword 支持全部三代格式;调用方可用 NeedsRehash 判定
// 登录成功后是否需要用新算法重哈希写回(登录升级路径见 api.adminLogin)。
package settings
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"strconv"
"strings"
"golang.org/x/crypto/bcrypt"
)
// bcryptCost bcrypt 工作因子:122026 年桌面 CPU 单次校验约 100-250ms
// 离线爆破成本相比单轮 SHA256 提升数个数量级)。
const bcryptCost = 12
// bcryptMaxLen bcrypt 算法只取前 72 字节;超长输入统一截断,
// 避免 GenerateFromPassword/CompareHashAndPassword 对 >72 字节返回错误。
const bcryptMaxLen = 72
func bcryptBytes(password string) []byte {
b := []byte(password)
if len(b) > bcryptMaxLen {
b = b[:bcryptMaxLen]
}
return b
}
// HashPassword 生成 bcrypt$<hash> 格式密码哈希(<hash> 为 bcrypt 原生串
// `$2a$<cost>$<salt><hash>`cost 内嵌于哈希串中)。
func HashPassword(password string) string {
sum, err := bcrypt.GenerateFromPassword(bcryptBytes(password), bcryptCost)
if err != nil {
// 截断后仅剩非法 cost 等实现级错误:确定性失败优于弱哈希回落
panic("settings: bcrypt 哈希失败: " + err.Error())
}
return "bcrypt$" + string(sum)
}
// VerifyPassword 校验密码:支持 bcrypt$、sha256$salt$hash 与旧版明文三种格式。
func VerifyPassword(password, hashed string) bool {
if hashed == "" {
return false
}
switch {
case strings.HasPrefix(hashed, "bcrypt$"):
return bcrypt.CompareHashAndPassword([]byte(hashed[len("bcrypt$"):]), bcryptBytes(password)) == nil
case strings.HasPrefix(hashed, "sha256$"):
parts := strings.Split(hashed, "$")
if len(parts) != 3 {
return false
}
salt, stored := parts[1], parts[2]
sum := sha256.Sum256([]byte(salt + password))
return hmac.Equal([]byte(hex.EncodeToString(sum[:])), []byte(stored))
}
// 旧版明文比较(兼容迁移)
return hmac.Equal([]byte(password), []byte(hashed))
}
// NeedsRehash 判断哈希是否需要升级为当前算法/成本(登录成功后判定,透明迁移)。
// sha256 与明文一律 truebcrypt 成本低于当前 bcryptCost 时 true。
func NeedsRehash(hashed string) bool {
if !strings.HasPrefix(hashed, "bcrypt$") {
return true
}
// bcrypt 原生串格式:$2a$<cost>$<salt><hash>
parts := strings.Split(hashed[len("bcrypt$"):], "$")
if len(parts) < 4 {
return true
}
cost, err := strconv.Atoi(parts[2])
if err != nil {
return true
}
return cost < bcryptCost
}
// IsPasswordHashed 判断是否为受支持的哈希格式(bcrypt / sha256)。
func IsPasswordHashed(s string) bool {
return strings.HasPrefix(s, "bcrypt$") || strings.HasPrefix(s, "sha256$")
}
// GenerateJWTSecret 生成 64 字符十六进制随机密钥。
func GenerateJWTSecret() string {
b := make([]byte, 32)
if _, err := rand.Read(b); err != nil {
panic("settings: crypto/rand 不可用: " + err.Error())
}
return hex.EncodeToString(b)
}
+38
View File
@@ -0,0 +1,38 @@
package settings
import "testing"
func TestHashPasswordRoundTrip(t *testing.T) {
h := HashPassword("s3cret-密码")
if !IsPasswordHashed(h) {
t.Fatalf("哈希格式不对: %s", h)
}
if !VerifyPassword("s3cret-密码", h) {
t.Fatal("正确密码校验失败")
}
if VerifyPassword("wrong", h) {
t.Fatal("错误密码竟通过校验")
}
if h == HashPassword("s3cret-密码") {
t.Fatal("盐值未随机化")
}
}
func TestVerifyLegacyPlaintext(t *testing.T) {
if !VerifyPassword("FileCodeBox2023", "FileCodeBox2023") {
t.Fatal("旧版明文兼容校验失败")
}
if VerifyPassword("nope", "FileCodeBox2023") {
t.Fatal("明文比较不应放行其他密码")
}
}
func TestGenerateJWTSecretLength(t *testing.T) {
s := GenerateJWTSecret()
if len(s) < 32 {
t.Fatalf("密钥太短: %d", len(s))
}
if s == GenerateJWTSecret() {
t.Fatal("密钥未随机化")
}
}
+258
View File
@@ -0,0 +1,258 @@
// Package settings — sanitize.go:受控 HTML 白名单净化(安全审计 L7)。
//
// notify_content 设计上「允许 <a> 等受控 HTML」,此前由管理端任意写入并经
// 前端 v-html 直出——管理员账号一旦被盗即可对全站访客注入脚本。
// 本净化器只保留纯文本与 <a href="http(s)|/|#">,其余标签连同其内层内容
// 一并丢弃(不做 HTML 转义输出,避免脚本字面量进入页面 DOM),
// 在公开配置读取与保存两处调用(双保险,覆盖历史存量数据)。
package settings
import (
"strings"
)
// dropContentTags 标签内部内容也一并丢弃的危险标签(script/style 等)。
var dropContentTags = map[string]bool{
"script": true, "style": true, "iframe": true, "object": true, "embed": true,
"title": true, "textarea": true, "noscript": true, "template": true,
"svg": true, "math": true, "xmp": true, "noembed": true, "noframes": true,
}
// SanitizeInlineHTML 白名单净化内联 HTML
// - <script>/<style>/<iframe> 等危险标签连同内部内容整体丢弃;
// - 其他非 <a> 标签仅丢弃标签本身、保留其内层文本(如 <b>加粗</b> → 加粗);
// - <a> 仅保留 href 属性,且值必须以 http://、https://、/ 或 # 开头;
// - HTML 注释(<!-- -->)丢弃,未闭合的危险标签丢弃其后全部内容;
// - 文本片段原样保留(不含 '<',渲染时为安全文本节点)。
func SanitizeInlineHTML(input string) string {
if input == "" {
return ""
}
var b strings.Builder
b.Grow(len(input))
i := 0
pendingAnchor := false
writeClose := func() {
if pendingAnchor {
b.WriteString("</a>")
pendingAnchor = false
}
}
for i < len(input) {
lt := strings.IndexByte(input[i:], '<')
if lt < 0 {
b.WriteString(input[i:])
break
}
b.WriteString(input[i : i+lt])
rest := input[i+lt:]
// 注释:整体丢弃
if strings.HasPrefix(rest, "<!--") {
end := strings.Index(rest, "-->")
if end < 0 {
break // 未闭合注释:丢弃剩余全部
}
i += lt + end + 3
continue
}
end := findTagEnd(rest)
if end < 0 {
break // 未闭合标签:丢弃剩余全部(不当作文本,防 < 绕过)
}
rawTag := rest[:end+1] // 形如 "<a href=..>"、"</div>"、"<img .../>"
name, closing, _ := parseTagName(rawTag)
if name != "" && !closing {
if dropContentTags[name] {
// 危险标签:连内层跳到对应闭合标签;无闭合(如 <script> 到结尾)则全丢
closeIdx := findClosingTag(input, i+lt+end+1, name)
if closeIdx < 0 {
writeClose()
return b.String()
}
i = closeIdx
continue
}
if name == "a" {
writeClose()
if href, ok := parseAllowedAnchor(rawTag); ok {
b.WriteString(`<a href="` + escapeAttr(href) + `">`)
pendingAnchor = true
}
// href 非法的 <a>:标签丢弃,但内层文本仍保留
}
// 其余开标签:丢弃标签本身,保留内层文本
i += lt + end + 1
continue
}
if name != "" && closing && name == "a" {
writeClose() // 仅在存在未闭合的合法 <a> 时输出
}
// 其余闭标签:丢弃
i += lt + end + 1
}
writeClose()
return b.String()
}
// findTagEnd 返回标签结束 '>' 的下标(跳过引号内的 '>',如 href="a<b">);未找到返回 -1。
func findTagEnd(s string) int {
inQuote := byte(0)
for i := 0; i < len(s); i++ {
c := s[i]
if inQuote != 0 {
if c == inQuote {
inQuote = 0
}
continue
}
switch c {
case '"', '\'':
inQuote = c
case '>':
return i
}
}
return -1
}
// parseTagName 解析标签名:返回 (小写名, 是否闭合标签, 是否自闭合 "/>")。
func parseTagName(tag string) (name string, closing, selfClosing bool) {
if len(tag) < 3 || tag[0] != '<' || tag[len(tag)-1] != '>' {
return "", false, false
}
inner := tag[1 : len(tag)-1]
if strings.HasSuffix(inner, "/") {
selfClosing = true
inner = inner[:len(inner)-1]
}
if strings.HasPrefix(inner, "/") {
closing = true
inner = inner[1:]
}
end := 0
for end < len(inner) {
r := inner[end]
if r == ' ' || r == '\t' || r == '\n' || r == '\r' {
break
}
end++
}
if end == 0 {
return "", closing, selfClosing
}
return strings.ToLower(inner[:end]), closing, selfClosing
}
// findClosingTag 从 from 开始查找 </name>,返回闭合标签结束位置(不含);找不到返回 -1。
func findClosingTag(s string, from int, name string) int {
needle := "</" + name
lower := strings.ToLower(s)
pos := from
for {
idx := strings.Index(lower[pos:], needle)
if idx < 0 {
return -1
}
at := pos + idx
after := at + len(needle)
if after < len(s) {
r := lower[after]
if r != '>' && r != ' ' && r != '\t' && r != '\n' && r != '\r' && r != '/' {
pos = after
continue // 形如 </scriptx> 的伪闭合,继续找
}
}
end := strings.IndexByte(s[after:], '>')
if end < 0 {
return -1
}
return after + end + 1
}
}
// parseAllowedAnchor 解析 <a ...> 标签:仅当 href 合法时返回 (href, true)。
func parseAllowedAnchor(tag string) (string, bool) {
inner := tag[1 : len(tag)-1]
// 标签名
nameEnd := 0
for nameEnd < len(inner) {
r := inner[nameEnd]
if r == ' ' || r == '\t' || r == '\n' || r == '\r' {
break
}
nameEnd++
}
href, found := scanAttr(inner[nameEnd:], "href")
if !found {
return "", false // 无 href 的 <a> 不放行(避免依赖默认行为)
}
href = strings.TrimSpace(href)
lower := strings.ToLower(href)
if !(strings.HasPrefix(lower, "http://") || strings.HasPrefix(lower, "https://") ||
strings.HasPrefix(href, "/") || strings.HasPrefix(href, "#")) {
return "", false // javascript:/data: 等一律拒绝
}
return href, true
}
// scanAttr 扫描属性串中的目标属性(支持双引号/单引号/无引号值)。
func scanAttr(s, name string) (string, bool) {
lower := strings.ToLower(s)
want := strings.ToLower(name)
for i := 0; i < len(lower); {
// 跳过空白
for i < len(lower) && (lower[i] == ' ' || lower[i] == '\t' || lower[i] == '\n' || lower[i] == '\r') {
i++
}
if i >= len(lower) {
break
}
// 属性名
start := i
for i < len(lower) && lower[i] != '=' && lower[i] != ' ' && lower[i] != '\t' && lower[i] != '\n' && lower[i] != '\r' {
i++
}
attrName := lower[start:i]
// 跳过空白
for i < len(lower) && (lower[i] == ' ' || lower[i] == '\t' || lower[i] == '\n' || lower[i] == '\r') {
i++
}
if i < len(lower) && lower[i] == '=' {
i++ // 跳过 '='
for i < len(lower) && (lower[i] == ' ' || lower[i] == '\t' || lower[i] == '\n' || lower[i] == '\r') {
i++
}
var val string
if i < len(lower) && (lower[i] == '"' || lower[i] == '\'') {
q := lower[i]
i++
vs := i
for i < len(lower) && lower[i] != q {
i++
}
val = s[vs:i]
if i < len(lower) {
i++ // 跳过闭合引号
}
} else {
vs := i
for i < len(lower) && lower[i] != ' ' && lower[i] != '\t' && lower[i] != '\n' && lower[i] != '\r' {
i++
}
val = s[vs:i]
}
if attrName == want {
return val, true
}
} else if attrName == want {
return "", true // 布尔属性:存在即命中(值空,调用方按非法处理)
}
}
return "", false
}
// escapeAttr HTML 属性转义。
func escapeAttr(s string) string {
r := strings.NewReplacer("&", "&amp;", "<", "&lt;", ">", "&gt;", `"`, "&#34;", "'", "&#39;")
return r.Replace(s)
}
+44
View File
@@ -0,0 +1,44 @@
// sanitize_test.go — SanitizeInlineHTML 单测(安全审计 L7)。
package settings
import "testing"
func TestSanitizeInlineHTML(t *testing.T) {
cases := []struct {
name string
in string
want string
}{
{"空串", "", ""},
{"纯文本保留", "欢迎使用文件快传", "欢迎使用文件快传"},
{"合法链接保留", `<a href="https://example.com">官网</a>`, `<a href="https://example.com">官网</a>`},
{"相对路径链接", `<a href="/docs">文档</a>`, `<a href="/docs">文档</a>`},
{"锚点链接", `<a href="#top">顶部</a>`, `<a href="#top">顶部</a>`},
{"script 整体丢弃", `hello<script>alert(1)</script>world`, "helloworld"},
{"img 丢弃保留文本", `a<img src=x onerror=alert(1)>b`, "ab"},
{"javascript href 拒绝", `<a href="javascript:alert(1)">x</a>`, "x"},
{"data href 拒绝", `<a href="data:text/html,<script>">x</a>`, "x"},
{"事件属性不透传", `<a href="/x" onclick="evil()">y</a>`, `<a href="/x">y</a>`},
{"注释丢弃", `a<!-- secret -->b`, "ab"},
{"未闭合标签丢弃剩余", `ok<script>alert(1)`, "ok"},
{"iframe 丢弃", `<iframe src="//evil"></iframe>text`, "text"},
{"样式标签丢弃", `<style>*{}</style>plain`, "plain"},
{"嵌套危险标签", `<div onclick=e><b>bold</b></div>`, "bold"},
{"大小写标签", `<A HREF="https://e.com">L</A>`, `<a href="https://e.com">L</a>`},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := SanitizeInlineHTML(tc.in); got != tc.want {
t.Fatalf("SanitizeInlineHTML(%q) = %q, want %q", tc.in, got, tc.want)
}
})
}
}
func TestSanitizeInlineHTMLNoScriptContent(t *testing.T) {
// script 内部文本也必须丢弃(不做 HTML 转义输出,避免 alert 字样进入页面 DOM)
got := SanitizeInlineHTML(`<script>var x = "</b>"; alert(1)</script>fine`)
if got != "fine" {
t.Fatalf("script 内容应整体丢弃, got %q", got)
}
}
+87
View File
@@ -0,0 +1,87 @@
// Package settings — schema.gov2 配置键 schema 常量与元数据表。
//
// 键名常量的单一事实来源在 internal/config/schema.godefaults() 需引用);
// 本文件 re-export 供 API/管理层使用,并提供「键名/类型/默认值」全量表,
// 供管理端设置页与文档生成(t4)对齐。新增键必须同步:
// 1. config/schema.go 键名与边界常量
// 2. config/config.go defaults() 默认值
// 3. 本文件 KVSchema() 元数据行
// 4. schema 同步测试(config schema_test / settings schema_test
package settings
import "filecodebox/internal/config"
// —— 键名 re-export(与 config 包保持同一字符串,避免魔法值散落)——
const (
// 需求 ① 背景图
KeyBackground = config.KeyBackground
KeyBackgroundURL = config.KeyBackgroundURL
// 需求 ② 页脚
KeyFooterText = config.KeyFooterText
KeyFooterBeian = config.KeyFooterBeian
// 需求 ③ 系统通知
KeyNotifyEnabled = config.KeyNotifyEnabled
KeyNotifyTitle = config.KeyNotifyTitle
KeyNotifyContent = config.KeyNotifyContent
// 需求 ④ 保存策略与上传频率限制
KeyMaxSaveSeconds = config.KeyMaxSaveSeconds
KeyMaxSaveCount = config.KeyMaxSaveCount
KeyExpireStyle = config.KeyExpireStyle
KeyUploadCount = config.KeyUploadCount
KeyUploadMinute = config.KeyUploadMinute
// 需求 ④⑩ 存储策略
KeyUploadSize = config.KeyUploadSize
KeyMaxFileSize = config.KeyMaxFileSize
KeyAllowedTypes = config.KeyAllowedTypes
KeyStorageLimit = config.KeyStorageLimit
KeyOpenUpload = config.KeyOpenUpload
// v3 存储引擎
KeyStorageEngine = config.KeyStorageEngine
)
// —— 取值边界 re-export ——
const (
MaxSaveSecondsMax = config.MaxSaveSecondsMax // 最长保存秒数上限(365 天)
MaxSaveCountMax = config.MaxSaveCountMax // 保存次数上限
MaxFileSizeMax = config.MaxFileSizeMax // 单文件大小上限(10 GiB
BackgroundURLMaxLen = config.BackgroundURLMaxLen // 背景图 URL 长度上限
FooterTextMaxLen = config.FooterTextMaxLen // 页脚内容长度上限
FooterBeianMaxLen = config.FooterBeianMaxLen // 备案号长度上限
NotifyTitleMaxLen = config.NotifyTitleMaxLen // 通知标题长度上限
NotifyContentMaxLen = config.NotifyContentMaxLen // 通知内容长度上限
)
// 敏感键:不允许出现在管理端 config get 下发/前端可见集合中(双模式下一致生效)。
// v3:引擎凭据(webdav_password/s3_secret_access_key/aws_session_token)加入敏感集——
// 管理端 get 返回掩码占位,update 时空串/掩码=不修改;公开 config 永不下发。
var SensitiveKeys = []string{
"admin_token", "jwt_secret",
"webdav_password", "s3_secret_access_key", "aws_session_token",
}
// SensitiveMaskValue 敏感键掩码占位(管理端 get 展示用)。
const SensitiveMaskValue = "******"
// IsSensitiveKey 判断键是否为敏感键(config get 必须屏蔽)。
func IsSensitiveKey(key string) bool {
for _, k := range SensitiveKeys {
if k == key {
return true
}
}
return false
}
// KVSchema 返回 v2 全量配置键元数据(键名/类型/默认值/边界/说明)。
// 默认值必须与 config defaults() 一致(schema 同步测试保证)。
func KVSchema() []config.KVSchemaEntry { return config.KVSchema() }
// KVSchemaByKey 以键名为索引查看 schema;未知键返回 nil。
func KVSchemaByKey(key string) *config.KVSchemaEntry {
for i := range config.KVSchema() {
if config.KVSchema()[i].Key == key {
return &config.KVSchema()[i]
}
}
return nil
}
+172
View File
@@ -0,0 +1,172 @@
// Package settings 提供数据库 settings KV 的运行时读写:
// envFCB_*)提供基线,DB KV 覆盖可变项;管理端修改后立即生效。
package settings
import (
"context"
"encoding/json"
"errors"
"log"
"sync"
"time"
"gorm.io/gorm"
"filecodebox/internal/config"
"filecodebox/internal/model"
)
// settingsKey 数据库中的配置键(对齐参考实现)。
const settingsKey = "settings"
// Manager 设置管理器:线程安全,缓存 KV 覆盖到内存。
type Manager struct {
db *gorm.DB
cfg *config.Config
mu sync.RWMutex
secret string // jwt_secret(频繁使用,单独缓存)
initPwd string // admin_token 哈希(频繁使用,单独缓存)
}
// NewManager 构造设置管理器并加载 DB KV。
// ensure 默认配置行(首次启动时写入 settings 键)。
func NewManager(ctx context.Context, db *gorm.DB, cfg *config.Config) (*Manager, error) {
m := &Manager{db: db, cfg: cfg}
if err := m.ensureSettingsRow(ctx); err != nil {
return nil, err
}
if err := m.Reload(ctx); err != nil {
return nil, err
}
return m, nil
}
// ensureSettingsRow 首次启动时把默认安全配置写入 KV(对齐 ensure_settings_row)。
func (m *Manager) ensureSettingsRow(ctx context.Context) error {
var row model.KeyValue
err := m.db.WithContext(ctx).Where(model.KeyValue{Key: settingsKey}).First(&row).Error
if err == nil {
return nil
}
// 双方言:必须用 errors.Is 判定(sqlite 驱动错误链与字符串消息与 postgres 不同)
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
// 不存在:写入初始配置(不含 admin_token/jwt_secret,保持未初始化状态)
initial := map[string]any{}
raw, _ := json.Marshal(initial)
row = model.KeyValue{Key: settingsKey, Value: strPtr(string(raw))}
if err := m.db.WithContext(ctx).Create(&row).Error; err != nil {
return err
}
log.Println("[settings] 系统尚未初始化,请在浏览器中打开站点并完成管理员密码设置")
return nil
}
// Reload 从数据库加载 settings KV 并覆盖到运行时配置。
func (m *Manager) Reload(ctx context.Context) error {
var row model.KeyValue
err := m.db.WithContext(ctx).Where(model.KeyValue{Key: settingsKey}).First(&row).Error
if err != nil {
// 行不存在时保持现有覆盖(双方言:errors.Is 判定)
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return err
}
kv := map[string]any{}
if row.Value != nil && *row.Value != "" {
if err := json.Unmarshal([]byte(*row.Value), &kv); err != nil {
log.Printf("[settings] settings KV 解析失败: %v", err)
}
}
// 内部键不允许通过 KV 覆盖(_ 开头)
safe := map[string]any{}
for k, v := range kv {
if len(k) > 0 && k[0] == '_' {
continue
}
safe[k] = v
}
m.mu.Lock()
m.cfg.ApplyKV(safe)
m.secret, _ = safe["jwt_secret"].(string)
m.initPwd, _ = safe["admin_token"].(string)
m.mu.Unlock()
return nil
}
// Get 返回当前配置(只读使用;不要修改返回值)。
func (m *Manager) Get() *config.Config { return m.cfg }
// SecretProvider 返回 jwt_secret 读取函数(JWT 中间件用)。
func (m *Manager) SecretProvider() func() string {
return func() string {
m.mu.RLock()
defer m.mu.RUnlock()
return m.secret
}
}
// IsInitialized 系统是否已完成初始化(管理员密码已设置且非默认密码)。
func (m *Manager) IsInitialized() bool {
m.mu.RLock()
defer m.mu.RUnlock()
if m.initPwd == "" {
return false
}
// 旧版默认密码视为未初始化(对齐 LEGACY_DEFAULT_ADMIN_TOKEN 检查)
return !verifyLegacyDefault(m.initPwd)
}
// legacyDefaultToken 参考实现的旧默认管理员密码。
const legacyDefaultToken = "FileCodeBox2023"
// verifyLegacyDefault 检查哈希是否对应旧默认密码。
func verifyLegacyDefault(hashed string) bool {
if hashed == "" {
return false
}
return VerifyPassword(legacyDefaultToken, hashed)
}
// UpdateKV 合并更新 settings KV(管理端保存配置)。
func (m *Manager) UpdateKV(ctx context.Context, patch map[string]any) error {
m.mu.Lock()
defer m.mu.Unlock()
// 读现有值
var row model.KeyValue
err := m.db.WithContext(ctx).Where(model.KeyValue{Key: settingsKey}).First(&row).Error
kv := map[string]any{}
if err == nil && row.Value != nil {
_ = json.Unmarshal([]byte(*row.Value), &kv)
}
for k, v := range patch {
if len(k) > 0 && k[0] == '_' {
continue
}
kv[k] = v
}
raw, err := json.Marshal(kv)
if err != nil {
return err
}
if err == nil && row.ID > 0 {
row.Value = strPtr(string(raw))
return m.db.WithContext(ctx).Model(&row).Update("value", row.Value).Error
}
row = model.KeyValue{Key: settingsKey, Value: strPtr(string(raw))}
return m.db.WithContext(ctx).Create(&row).Error
}
// SystemStart 记录系统启动时间(对齐 sys_start 键)。
func (m *Manager) SystemStart(ctx context.Context) {
now := time.Now().UnixMilli()
raw, _ := json.Marshal(now)
_ = m.db.WithContext(ctx).Where(model.KeyValue{Key: "sys_start"}).
Assign(model.KeyValue{Value: strPtr(string(raw))}).
FirstOrCreate(&model.KeyValue{Key: "sys_start"}).Error
}
func strPtr(s string) *string { return &s }
+10
View File
@@ -0,0 +1,10 @@
package storage
import "errors"
// ErrRangeNotSatisfiable 请求的字节范围超出文件大小(对应 HTTP 416)。
// 属于增量错误定义,不改动 interface.go 的既有签名。
var ErrRangeNotSatisfiable = errors.New("storage: 请求范围超出文件大小")
// ErrHashMismatch 分片哈希校验失败(合并分片时检测到数据损坏)。
var ErrHashMismatch = errors.New("storage: 分片哈希校验失败")
+26
View File
@@ -0,0 +1,26 @@
package storage
import (
"context"
"fmt"
)
// Factory 按配置构造存储引擎。由 go-storage 提供 New* 实现后接入。
// 这里提供注册表模式:各引擎实现注册自己的构造函数,main.go 按名称选择。
type Factory func(ctx context.Context) (Storage, error)
var registry = map[string]Factory{}
// RegisterEngine 注册引擎构造函数(init 时调用,名称:local|s3|webdav)。
func RegisterEngine(name string, f Factory) {
registry[name] = f
}
// NewEngine 按名称构造引擎。
func NewEngine(ctx context.Context, name string) (Storage, error) {
f, ok := registry[name]
if !ok {
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav", name)
}
return f(ctx)
}
+111
View File
@@ -0,0 +1,111 @@
// Package storage 定义存储引擎统一契约。
//
// 本文件是 go-storage 并行开发的接口契约:签名一经定义不再改动。
// 三种引擎(local/s3/webdav)都要实现该接口;工厂按 FCB_STORAGE_ENGINE 选择。
package storage
import (
"context"
"errors"
"io"
)
// 错误定义:实现方应返回这些哨兵错误(可用 %w 包装),便于 API 层映射 HTTP 状态码。
var (
// ErrNotFound 文件不存在(HTTP 404)。
ErrNotFound = errors.New("storage: 文件不存在")
// ErrInvalidPath 非法路径(路径穿越等,HTTP 400)。
ErrInvalidPath = errors.New("storage: 非法文件路径")
// ErrUnavailable 存储服务不可用(连接失败等,HTTP 503)。
ErrUnavailable = errors.New("storage: 存储服务不可用")
)
// FileMeta 文件元信息(大小等)。
type FileMeta struct {
Size int64 // 字节数
ContentType string // MIME 类型,可为空
AcceptRanges bool // 是否支持 Range 请求
}
// Download 流式下载句柄。调用方负责 Close。
type Download struct {
// ReadCloser 文件内容流(已按 Range 重定位)。
io.ReadCloser
// Meta 文件元信息。
Meta FileMeta
// Start 当前流的起始字节偏移(Range 请求时为 rangeStart)。
Start int64
// End 流的结束字节偏移(含);未知为 -1。
End int64
// Total 文件总大小(字节);未知为 -1。
Total int64
}
// Range 字节范围(对齐 HTTP Range 语义)。
// nil 指针表示完整文件。
type Range struct {
Start int64 // 起始字节(含)
End int64 // 结束字节(含);-1 表示到文件末尾
}
// Storage 存储引擎统一接口。
//
// 约定:
// - savePath 为存储侧相对路径(引擎内部负责安全解析,拒绝 .. 穿越);
// - 所有方法必须是并发安全的;
// - 实现方遇到不可恢复错误时返回本包哨兵错误(或用 %w 包装)。
type Storage interface {
// SaveFile 流式保存文件:r 读取到 EOF 即完成,返回实际写入字节数。
// 引擎必须按 256KB 级别分块读取,不得将整个文件读入内存。
SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error)
// DeleteFile 删除文件;文件不存在时返回 ErrNotFound 或 nil 均可接受。
DeleteFile(ctx context.Context, savePath string) error
// Open 以下载模式打开文件,支持 HTTP Range 请求语义:
// - rng 为 nil:返回完整文件流(Start=0End=Total-1);
// - rng 非 nil:返回 [Start, End] 区间流。
// 引擎应尽量透传 RangeWebDAV/S3)或按块 seeklocal)。
Open(ctx context.Context, savePath string, rng *Range) (*Download, error)
// Stat 获取文件元信息;不存在返回 ErrNotFound。
Stat(ctx context.Context, savePath string) (*FileMeta, error)
// SaveChunk 保存一个分片到临时区(upload_id 隔离),返回分片字节数。
SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error)
// MergeChunks 按索引 0..total-1 有序合并分片并落为正式文件。
// verifyHash 为nil 时不校验;否则为分片 SHA256 校验函数(输入索引,输出期望哈希,空串表示跳过)。
// 返回 (最终文件大小, 整个文件 SHA256)。
MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error)
// CleanChunks 清理分片临时区;不存在时静默成功。
CleanChunks(ctx context.Context, uploadID string, savePath string) error
// FileExists 检查文件是否存在。
FileExists(ctx context.Context, savePath string) (bool, error)
// HeadMeta 读取对象元信息与头部字节(可选能力,供直传 confirm 校验实际
// 大小与内容;不支持时返回 ErrNotSupported)。
// meta 允许为 nil(仅取头部);head 为对象前 headBytes 字节(不足时取实际长度)。
HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error)
// PresignGetURL 生成限时直链(下载);不支持直链的引擎返回 ErrNotSupported。
PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error)
// PresignPutURL 生成限时直传(上传)URL;不支持直传的引擎返回 ErrNotSupported。
PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error)
// HealthCheck 引擎健康检查(启动时与 /health 使用)。
HealthCheck(ctx context.Context) error
}
// ErrNotSupported 当前引擎不支持该能力(如本地引擎不支持预签名)。
var ErrNotSupported = errors.New("storage: 当前引擎不支持该操作")
// ChunkPath 返回分片临时路径(约定统一为 <dir>/chunks/<upload_id>/<index>.part)。
// 引擎可使用 ChunkDir 拼接自身路径。
type PathBuilder interface {
// ChunkDir 分片临时目录(相对 savePath 所在目录)。
ChunkDir(savePath, uploadID string) string
}
+449
View File
@@ -0,0 +1,449 @@
package storage
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"mime"
"os"
"path/filepath"
"strings"
"sync/atomic"
)
// 每次读写使用的缓冲大小:256KB,对齐参考实现 SystemFileStorage.chunk_size。
const localChunkSize = 256 * 1024
// LocalStorage 本地文件系统引擎。
//
// 相比参考实现(SystemFileStorage)的改进:
// - 双重路径防护:清洗相对路径 + 根目录前缀校验 + 符号链接逃逸校验;
// - 全部落盘走「临时文件 + fsync + 原子重命名」,断电/中断不产生半截文件;
// - 下载使用 io.NewSectionReader 支持任意 Range,无需整文件读入内存。
type LocalStorage struct {
// root 存储根目录(绝对路径)。
root string
// rootReal 经符号链接解析后的真实根目录,用于逃逸校验。
rootReal string
}
// NewLocalStorage 构造本地引擎。root 为空时使用系统临时目录下的 filecodebox_storage。
func NewLocalStorage(root string) (*LocalStorage, error) {
if strings.TrimSpace(root) == "" {
root = filepath.Join(os.TempDir(), "filecodebox_storage")
}
abs, err := filepath.Abs(root)
if err != nil {
return nil, fmt.Errorf("storage/local: 解析根目录失败: %w", err)
}
if err := os.MkdirAll(abs, 0o755); err != nil {
return nil, fmt.Errorf("storage/local: 创建根目录失败: %w", err)
}
real := abs
if resolved, err := filepath.EvalSymlinks(abs); err == nil {
real = resolved
}
return &LocalStorage{root: abs, rootReal: real}, nil
}
func init() {
RegisterEngine("local", func(ctx context.Context) (Storage, error) {
return NewLocalStorage(engineOptions.Local.Root)
})
}
// withinRoot 判断路径 p 是否位于 root 内(含 root 本身)。
func withinRoot(p, root string) bool {
p = filepath.Clean(p)
root = filepath.Clean(root)
if p == root {
return true
}
return strings.HasPrefix(p, root+string(os.PathSeparator))
}
// absPath 将存储侧相对路径解析为根目录内的绝对路径。
// 任何路径穿越或符号链接逃逸都会返回 ErrInvalidPath。
func (l *LocalStorage) absPath(savePath string) (string, error) {
cleaned, ok := SanitizePath(savePath)
if !ok {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
full := filepath.Join(l.root, filepath.FromSlash(cleaned))
if !withinRoot(full, l.root) {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
// 符号链接逃逸校验:文件已存在时解析真实路径;不存在时校验最深已存在的父目录。
if real, err := filepath.EvalSymlinks(full); err == nil {
if !withinRoot(real, l.rootReal) {
return "", fmt.Errorf("%w: 符号链接逃逸 %q", ErrInvalidPath, savePath)
}
} else {
dir := filepath.Dir(full)
if realDir, err := filepath.EvalSymlinks(dir); err == nil && !withinRoot(realDir, l.rootReal) {
return "", fmt.Errorf("%w: 符号链接逃逸 %q", ErrInvalidPath, savePath)
}
}
return full, nil
}
// SaveFile 流式保存:256KB 分块读取写入临时文件,fsync 后原子重命名。
func (l *LocalStorage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
full, err := l.absPath(savePath)
if err != nil {
return 0, err
}
src := &countingReader{r: r}
if err := writeFileAtomic(full, src); err != nil {
return src.count(), err
}
return src.count(), nil
}
// DeleteFile 删除文件;文件不存在时静默成功(对齐契约)。
func (l *LocalStorage) DeleteFile(ctx context.Context, savePath string) error {
full, err := l.absPath(savePath)
if err != nil {
return err
}
if err := os.Remove(full); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("storage/local: 删除失败: %w", err)
}
return nil
}
// Open 打开文件下载流;rng 非 nil 时用 SectionReader 实现 Range 语义。
func (l *LocalStorage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
full, err := l.absPath(savePath)
if err != nil {
return nil, err
}
f, err := os.Open(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, ErrNotFound
}
return nil, fmt.Errorf("storage/local: 打开文件失败: %w", err)
}
info, err := f.Stat()
if err != nil {
_ = f.Close()
return nil, fmt.Errorf("storage/local: 获取文件信息失败: %w", err)
}
if info.IsDir() {
_ = f.Close()
return nil, ErrNotFound
}
size := info.Size()
start, end := int64(0), size-1
if rng != nil {
if rng.Start < 0 || (rng.End != -1 && rng.End < rng.Start) {
_ = f.Close()
return nil, fmt.Errorf("%w: 非法 Range", ErrRangeNotSatisfiable)
}
if rng.Start >= size {
_ = f.Close()
return nil, ErrRangeNotSatisfiable
}
start = rng.Start
end = size - 1
if rng.End != -1 && rng.End < end {
end = rng.End
}
}
section := io.NewSectionReader(f, start, end-start+1)
dl := &Download{
ReadCloser: &fileSection{Reader: section, closer: f},
Start: start,
End: end,
Total: size,
Meta: FileMeta{
Size: size,
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
AcceptRanges: true,
},
}
if end < 0 { // 空文件:End 语义上等于 -1(未知),Total=0 已表达大小
dl.End = -1
}
return dl, nil
}
// fileSection 组合 SectionReader 与文件关闭器。
type fileSection struct {
io.Reader
closer io.Closer
}
func (f *fileSection) Close() error { return f.closer.Close() }
// Stat 获取文件元信息。
func (l *LocalStorage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
full, err := l.absPath(savePath)
if err != nil {
return nil, err
}
info, err := os.Stat(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, ErrNotFound
}
return nil, fmt.Errorf("storage/local: Stat 失败: %w", err)
}
if info.IsDir() {
return nil, ErrNotFound
}
return &FileMeta{
Size: info.Size(),
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
AcceptRanges: true,
}, nil
}
// SaveChunk 保存分片到 <父目录>/chunks/<uploadID>/<index>.part,原子写入。
func (l *LocalStorage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
// 先校验目标路径合法性,分片目录随合法路径派生。
if _, err := l.absPath(savePath); err != nil {
return 0, err
}
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, chunkIndex))
if !ok {
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, savePath)
}
full := filepath.Join(l.root, filepath.FromSlash(chunkRel))
if !withinRoot(full, l.root) {
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, chunkRel)
}
src := &countingReader{r: r}
if err := writeFileAtomic(full, src); err != nil {
return src.count(), err
}
return src.count(), nil
}
// MergeChunks 按索引 0..total-1 有序合并分片:
// - 逐分片流式拷贝到临时输出(边拷贝边计算整文件与分片 SHA256);
// - verifyHash 非 nil 时校验分片哈希(空串跳过);
// - 全部通过后 fsync + 原子重命名,并清理分片临时目录。
func (l *LocalStorage) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
if total <= 0 {
return 0, "", fmt.Errorf("storage/local: 非法分片总数 %d", total)
}
full, err := l.absPath(savePath)
if err != nil {
return 0, "", err
}
if err := os.MkdirAll(filepath.Dir(full), 0o755); err != nil {
return 0, "", fmt.Errorf("storage/local: 创建目标目录失败: %w", err)
}
tmp, err := os.CreateTemp(filepath.Dir(full), "."+filepath.Base(full)+".merging-*")
if err != nil {
return 0, "", fmt.Errorf("storage/local: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() {
_ = tmp.Close()
_ = os.Remove(tmpName) // 成功时已被重命名,删除静默失败
}()
totalHash := sha256.New()
buf := make([]byte, localChunkSize)
var size int64
for i := 0; i < total; i++ {
if err := ctx.Err(); err != nil {
return 0, "", err
}
partRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, i))
if !ok {
return 0, "", fmt.Errorf("%w: 分片 %d 路径非法", ErrInvalidPath, i)
}
partPath := filepath.Join(l.root, filepath.FromSlash(partRel))
in, err := os.Open(partPath)
if err != nil {
return 0, "", fmt.Errorf("storage/local: 分片 %d 不存在: %w", i, err)
}
chunkHash := sha256.New()
n, err := io.CopyBuffer(io.MultiWriter(tmp, totalHash, chunkHash), in, buf)
_ = in.Close()
if err != nil {
return 0, "", fmt.Errorf("storage/local: 读取分片 %d 失败: %w", i, err)
}
if verifyHash != nil {
expected, err := verifyHash(i)
if err != nil {
return 0, "", err
}
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
}
}
size += n
}
if err := tmp.Sync(); err != nil {
return 0, "", fmt.Errorf("storage/local: 落盘失败: %w", err)
}
if err := tmp.Close(); err != nil {
return 0, "", fmt.Errorf("storage/local: 关闭临时文件失败: %w", err)
}
if err := os.Rename(tmpName, full); err != nil {
return 0, "", fmt.Errorf("storage/local: 原子重命名失败: %w", err)
}
// 合并成功后清理分片临时目录(静默容错,不掩盖成功结果)。
_ = l.CleanChunks(ctx, uploadID, savePath)
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
}
// CleanChunks 清理分片临时目录;不存在时静默成功,并尝试移除空 chunks 父目录。
func (l *LocalStorage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
dirRel, ok := SanitizePath(chunkDirOf(savePath, uploadID))
if !ok {
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, savePath)
}
dir := filepath.Join(l.root, filepath.FromSlash(dirRel))
if !withinRoot(dir, l.root) {
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, dirRel)
}
if err := os.RemoveAll(dir); err != nil {
return fmt.Errorf("storage/local: 清理分片目录失败: %w", err)
}
// 父级 chunks 目录为空则一并清理(对齐参考实现)。
chunksParent := filepath.Dir(dir)
if entries, err := os.ReadDir(chunksParent); err == nil && len(entries) == 0 {
_ = os.Remove(chunksParent)
}
return nil
}
// chunkDirOf 返回分片目录(去掉文件名部分):<父目录>/chunks/<uploadID>。
func chunkDirOf(savePath, uploadID string) string {
cd := ChunkDir(savePath, uploadID)
// ChunkDir 返回 "<dir>/chunks/<uploadID>/<name>",去掉末段文件名即目录。
if idx := strings.LastIndex(cd, "/"); idx > 0 {
return cd[:idx]
}
return cd
}
// FileExists 检查文件是否存在;非法路径按不存在处理(对齐参考实现)。
func (l *LocalStorage) FileExists(ctx context.Context, savePath string) (bool, error) {
full, err := l.absPath(savePath)
if err != nil {
return false, nil
}
info, err := os.Stat(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return false, nil
}
return false, fmt.Errorf("storage/local: Stat 失败: %w", err)
}
return !info.IsDir(), nil
}
// HeadMeta 读取文件元信息与前 n 字节(本地引擎实现)。
func (l *LocalStorage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
full, err := l.absPath(savePath)
if err != nil {
return nil, nil, err
}
f, err := os.Open(full)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, nil, ErrNotFound
}
return nil, nil, fmt.Errorf("storage/local: 打开文件失败: %w", err)
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil, nil, fmt.Errorf("storage/local: Stat 失败: %w", err)
}
if info.IsDir() {
return nil, nil, ErrNotFound
}
head := make([]byte, headBytes)
n, _ := io.ReadFull(f, head)
return &FileMeta{
Size: info.Size(),
ContentType: mime.TypeByExtension(strings.ToLower(filepath.Ext(full))),
AcceptRanges: true,
}, head[:n], nil
}
// PresignGetURL 本地引擎不支持直链。
func (l *LocalStorage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// PresignPutURL 本地引擎不支持直传。
func (l *LocalStorage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// HealthCheck 健康检查:根目录可写(写入并删除探针文件)。
func (l *LocalStorage) HealthCheck(ctx context.Context) error {
if err := os.MkdirAll(l.root, 0o755); err != nil {
return fmt.Errorf("%w: 本地存储根目录不可创建: %v", ErrUnavailable, err)
}
probe := filepath.Join(l.root, ".health-probe")
if err := os.WriteFile(probe, []byte("ok"), 0o644); err != nil {
return fmt.Errorf("%w: 本地存储不可写: %v", ErrUnavailable, err)
}
_ = os.Remove(probe)
return nil
}
// writeFileAtomic 临时文件 + fsync + rename 的原子落盘。
func writeFileAtomic(dst string, src io.Reader) error {
dir := filepath.Dir(dst)
if err := os.MkdirAll(dir, 0o755); err != nil {
return fmt.Errorf("storage/local: 创建目录失败: %w", err)
}
tmp, err := os.CreateTemp(dir, "."+filepath.Base(dst)+".tmp-*")
if err != nil {
return fmt.Errorf("storage/local: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
cleanup := func() { _ = tmp.Close(); _ = os.Remove(tmpName) }
if _, err := io.CopyBuffer(tmp, src, make([]byte, localChunkSize)); err != nil {
cleanup()
return fmt.Errorf("storage/local: 写入失败: %w", err)
}
if err := tmp.Sync(); err != nil {
cleanup()
return fmt.Errorf("storage/local: fsync 失败: %w", err)
}
if err := tmp.Close(); err != nil {
_ = os.Remove(tmpName)
return fmt.Errorf("storage/local: 关闭临时文件失败: %w", err)
}
if err := os.Rename(tmpName, dst); err != nil {
_ = os.Remove(tmpName)
return fmt.Errorf("storage/local: 原子重命名失败: %w", err)
}
return nil
}
// countingReader 统计累计读取字节数(并发安全)。
type countingReader struct {
r io.Reader
n atomic.Int64
}
func (c *countingReader) Read(p []byte) (int, error) {
n, err := c.r.Read(p)
c.n.Add(int64(n))
return n, err
}
// count 返回累计字节数。
func (c *countingReader) count() int64 { return c.n.Load() }
// reset 归零计数(请求体重放时使用)。
func (c *countingReader) reset() { c.n.Store(0) }
// 接口编译期断言。
var _ Storage = (*LocalStorage)(nil)
+346
View File
@@ -0,0 +1,346 @@
package storage
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"io"
"os"
"path/filepath"
"strings"
"sync"
"testing"
)
// newTestLocal 构造以临时目录为根的本地引擎。
func newTestLocal(t *testing.T) *LocalStorage {
t.Helper()
st, err := NewLocalStorage(t.TempDir())
if err != nil {
t.Fatalf("NewLocalStorage: %v", err)
}
return st
}
func sha256Hex(b []byte) string {
sum := sha256.Sum256(b)
return hex.EncodeToString(sum[:])
}
// TestLocalSaveOpenRange 保存/Stat/完整与 Range 下载。
func TestLocalSaveOpenRange(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
data := []byte("hello filecodebox 本地引擎 0123456789")
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/测试文件.bin")
if err != nil {
t.Fatalf("SaveFile: %v", err)
}
if n != int64(len(data)) {
t.Fatalf("SaveFile n = %d, want %d", n, len(data))
}
// Stat
meta, err := st.Stat(ctx, "2025/08/测试文件.bin")
if err != nil {
t.Fatalf("Stat: %v", err)
}
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
t.Fatalf("Stat = %+v", meta)
}
// 完整下载:对齐 go-api 约定 Start=0、End=Total-1
dl, err := st.Open(ctx, "2025/08/测试文件.bin", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, err := io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read full: %v", err)
}
if !bytes.Equal(got, data) {
t.Fatalf("full content mismatch")
}
if dl.Start != 0 || dl.End != int64(len(data))-1 || dl.Total != int64(len(data)) {
t.Fatalf("full download offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// Range 下载 [2, 7]
dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 2, End: 7})
if err != nil {
t.Fatalf("Open range: %v", err)
}
got, err = io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read range: %v", err)
}
if !bytes.Equal(got, data[2:8]) {
t.Fatalf("range content = %q, want %q", got, data[2:8])
}
// Range end 越界自动钳制
dl, err = st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: 5, End: 99999})
if err != nil {
t.Fatalf("Open clamp range: %v", err)
}
got, _ = io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data[5:]) {
t.Fatalf("clamp range mismatch")
}
// 起点越界 → 416
if _, err := st.Open(ctx, "2025/08/测试文件.bin", &Range{Start: int64(len(data)) + 1, End: -1}); err == nil {
t.Fatalf("out-of-range start should fail")
} else if !strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
}
// 不存在 → 404
if _, err := st.Open(ctx, "no/such/file.bin", nil); err == nil || !strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("want ErrNotFound, got %v", err)
}
}
// TestLocalSaveFileAtomic 原子写:目录自动创建,无临时残留。
func TestLocalSaveFileAtomic(t *testing.T) {
ctx := context.Background()
dir := t.TempDir()
st2, _ := NewLocalStorage(dir)
big := bytes.Repeat([]byte("abc123"), 128*1024) // 768KB,跨多个 256KB 缓冲
if _, err := st2.SaveFile(ctx, bytes.NewReader(big), "a/b/c/big.bin"); err != nil {
t.Fatalf("SaveFile: %v", err)
}
// 无临时残留
entries, _ := os.ReadDir(filepath.Join(dir, "a", "b", "c"))
for _, e := range entries {
if strings.HasPrefix(e.Name(), ".") {
t.Fatalf("临时文件残留: %s", e.Name())
}
}
got, _ := os.ReadFile(filepath.Join(dir, "a", "b", "c", "big.bin"))
if !bytes.Equal(got, big) {
t.Fatalf("content mismatch")
}
}
// TestLocalTraversal 防路径穿越(含符号链接逃逸)。
func TestLocalTraversal(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
for _, p := range []string{"../escape.txt", "a/../../escape", "..", "/../x"} {
if _, err := st.SaveFile(ctx, strings.NewReader("x"), p); err == nil ||
!strings.Contains(err.Error(), ErrInvalidPath.Error()) {
t.Fatalf("SaveFile(%q) 应拒绝: %v", p, err)
}
if _, err := st.Open(ctx, p, nil); err == nil || !strings.Contains(err.Error(), ErrInvalidPath.Error()) {
t.Fatalf("Open(%q) 应拒绝: %v", p, err)
}
}
// 符号链接逃逸
outside := filepath.Join(t.TempDir(), "outside.txt")
if err := os.WriteFile(outside, []byte("secret"), 0o644); err != nil {
t.Fatal(err)
}
link := filepath.Join(st.root, "link.txt")
if err := os.Symlink(outside, link); err != nil {
t.Skipf("symlink 不可用: %v", err)
}
if _, err := st.Open(ctx, "link.txt", nil); err == nil ||
!strings.Contains(err.Error(), ErrInvalidPath.Error()) {
t.Fatalf("symlink escape 应拒绝: %v", err)
}
}
// TestLocalChunkLifecycle 分片保存/合并(索引有序 + SHA256 校验 + 清理临时目录)。
func TestLocalChunkLifecycle(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
savePath := "2025/09/chunked.bin"
uploadID := "upload-abc"
chunks := [][]byte{[]byte("AAAA"), []byte("BB"), []byte("CCCCCC")}
hashes := make([]string, len(chunks))
var total int64
for i, c := range chunks {
n, err := st.SaveChunk(ctx, uploadID, i, bytes.NewReader(c), savePath)
if err != nil {
t.Fatalf("SaveChunk %d: %v", i, err)
}
if n != int64(len(c)) {
t.Fatalf("SaveChunk %d n = %d", i, n)
}
hashes[i] = sha256Hex(c)
total += int64(len(c))
}
// 分片文件确实存在于临时目录
for i := range chunks {
exists, err := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, i))
if err != nil || !exists {
t.Fatalf("分片 %d 应存在: %v %v", i, exists, err)
}
}
size, fileHash, err := st.MergeChunks(ctx, uploadID, len(chunks), func(i int) (string, error) {
return hashes[i], nil
}, savePath)
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != total {
t.Fatalf("merged size = %d, want %d", size, total)
}
want := sha256Hex(bytes.Join(chunks, nil))
if fileHash != want {
t.Fatalf("file hash = %s, want %s", fileHash, want)
}
// 合并内容 = 按索引有序拼接
got, err := os.ReadFile(filepath.Join(st.root, filepath.FromSlash(savePath)))
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(got, bytes.Join(chunks, nil)) {
t.Fatalf("merged content mismatch")
}
// 合并后分片目录已清理
exists, _ := st.FileExists(ctx, ChunkPartPath(savePath, uploadID, 0))
if exists {
t.Fatalf("合并后分片应已清理")
}
if _, err := os.Stat(filepath.Join(st.root, "2025/09/chunks")); !os.IsNotExist(err) {
t.Fatalf("chunks 父目录应已清理: %v", err)
}
}
// TestLocalChunkHashMismatch 哈希不匹配 → 合并失败且不留输出文件。
func TestLocalChunkHashMismatch(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
savePath := "mismatch.bin"
if _, err := st.SaveChunk(ctx, "uid", 0, strings.NewReader("data"), savePath); err != nil {
t.Fatal(err)
}
_, _, err := st.MergeChunks(ctx, "uid", 1, func(i int) (string, error) {
return sha256Hex([]byte("WRONG")), nil
}, savePath)
if err == nil || !strings.Contains(err.Error(), ErrHashMismatch.Error()) {
t.Fatalf("want ErrHashMismatch, got %v", err)
}
exists, _ := st.FileExists(ctx, savePath)
if exists {
t.Fatalf("校验失败不应产出正式文件")
}
}
// TestLocalVerifyHashAbort verifyHash 返回 error → 合并中止。
func TestLocalVerifyHashAbort(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
if _, err := st.SaveChunk(ctx, "uid2", 0, strings.NewReader("data"), "v.bin"); err != nil {
t.Fatal(err)
}
boom := context.Canceled
if _, _, err := st.MergeChunks(ctx, "uid2", 1, func(i int) (string, error) {
return "", boom
}, "v.bin"); err == nil {
t.Fatalf("verifyHash 错误应向上传播")
}
}
// TestLocalDeleteExistsClean 删除/存在性/清理。
func TestLocalDeleteExistsClean(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
if _, err := st.SaveFile(ctx, strings.NewReader("x"), "d/f.bin"); err != nil {
t.Fatal(err)
}
if ok, _ := st.FileExists(ctx, "d/f.bin"); !ok {
t.Fatalf("文件应存在")
}
if err := st.DeleteFile(ctx, "d/f.bin"); err != nil {
t.Fatalf("DeleteFile: %v", err)
}
if ok, _ := st.FileExists(ctx, "d/f.bin"); ok {
t.Fatalf("文件应已删除")
}
// 删除不存在 → nil
if err := st.DeleteFile(ctx, "d/f.bin"); err != nil {
t.Fatalf("删除不存在应静默: %v", err)
}
// CleanChunks 幂等
if err := st.CleanChunks(ctx, "uid-x", "y.bin"); err != nil {
t.Fatalf("CleanChunks: %v", err)
}
}
// TestLocalPresignNotSupported 预签名 → ErrNotSupported。
func TestLocalPresignNotSupported(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
if _, err := st.PresignGetURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("PresignGetURL want ErrNotSupported, got %v", err)
}
if _, err := st.PresignPutURL(ctx, "a.bin", 60); err == nil || !strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("PresignPutURL want ErrNotSupported, got %v", err)
}
}
// TestLocalHealthCheck 健康检查。
func TestLocalHealthCheck(t *testing.T) {
st := newTestLocal(t)
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
}
// TestLocalConcurrent 并发安全冒烟。
func TestLocalConcurrent(t *testing.T) {
st := newTestLocal(t)
ctx := context.Background()
var wg sync.WaitGroup
for i := 0; i < 16; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
p := "conc/" + itoa(i) + ".bin"
if _, err := st.SaveFile(ctx, strings.NewReader(strings.Repeat("x", i+1)), p); err != nil {
t.Errorf("save %d: %v", i, err)
return
}
dl, err := st.Open(ctx, p, nil)
if err != nil {
t.Errorf("open %d: %v", i, err)
return
}
_, _ = io.ReadAll(dl)
_ = dl.Close()
}(i)
}
wg.Wait()
}
// TestLocalFactoryRegistry 工厂注册与构造。
func TestLocalFactoryRegistry(t *testing.T) {
prev := engineOptions.Local.Root
engineOptions.Local.Root = t.TempDir()
defer func() { engineOptions.Local.Root = prev }()
st, err := NewEngine(context.Background(), "local")
if err != nil {
t.Fatalf("NewEngine(local): %v", err)
}
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if _, err := NewEngine(context.Background(), "unknown"); err == nil {
t.Fatalf("未知引擎应报错")
}
}
+187
View File
@@ -0,0 +1,187 @@
package storage
import (
"context"
"fmt"
"io"
"sync"
)
// Manager 存储引擎管理器:实现 Storage 全接口并支持运行时热切换。
//
// v3 需求:管理后台可设置存储类型(local|s3|webdav)与各引擎参数,
// 保存后无需重启即生效。设计要点:
// - 读写/保存类操作全部委托到"当前引擎"(原子指针,无锁热路径);
// - Switch 先构建并健康检查新引擎,成功才替换指针,失败保持原引擎;
// - EngineOf 按名字取引擎实例(带缓存),供"按文件归属引擎取回旧文件"使用;
// - 管理端修改引擎参数后调用 Invalidate 使对应实例缓存失效,下次构建生效。
type Manager struct {
// build 构建指定引擎实例(由装配方注入:内部刷新全局 EngineOptions 后走工厂)。
build func(name string) (Storage, error)
mu sync.RWMutex
current Storage
curName string
cache map[string]Storage
}
// validEngines 合法引擎名(与 FCB_STORAGE_ENGINE 枚举一致)。
var validEngines = map[string]bool{"local": true, "s3": true, "webdav": true}
// ValidEngine 校验引擎名是否合法。
func ValidEngine(name string) bool { return validEngines[name] }
// NewManager 创建管理器:current 为启动时已构建的引擎(主装配流已做过健康检查)。
// build 注入构建函数(管理端切换/参数变更时使用,内部须串行——Manager 已加锁)。
func NewManager(name string, current Storage, build func(name string) (Storage, error)) *Manager {
return &Manager{
build: build,
current: current,
curName: name,
cache: map[string]Storage{name: current},
}
}
// —— Storage 接口委托(全部走当前引擎)——
// SaveFile 流式保存文件(委托当前引擎)。
func (m *Manager) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
return m.current.SaveFile(ctx, r, savePath)
}
// DeleteFile 删除文件(委托当前引擎)。
func (m *Manager) DeleteFile(ctx context.Context, savePath string) error {
return m.current.DeleteFile(ctx, savePath)
}
// Open 打开文件流(委托当前引擎;旧文件由 API 层先经 EngineOf 按归属引擎取)。
func (m *Manager) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
return m.current.Open(ctx, savePath, rng)
}
// Stat 文件元信息(委托当前引擎)。
func (m *Manager) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
return m.current.Stat(ctx, savePath)
}
// SaveChunk 保存分片(委托当前引擎)。
func (m *Manager) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
return m.current.SaveChunk(ctx, uploadID, chunkIndex, r, savePath)
}
// MergeChunks 合并分片(委托当前引擎)。
func (m *Manager) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
return m.current.MergeChunks(ctx, uploadID, total, verifyHash, savePath)
}
// CleanChunks 清理分片临时区(委托当前引擎)。
func (m *Manager) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
return m.current.CleanChunks(ctx, uploadID, savePath)
}
// FileExists 文件存在性(委托当前引擎)。
func (m *Manager) FileExists(ctx context.Context, savePath string) (bool, error) {
return m.current.FileExists(ctx, savePath)
}
// HeadMeta 元信息与头部字节(委托当前引擎;供直传 confirm 校验)。
func (m *Manager) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
return m.current.HeadMeta(ctx, savePath, headBytes)
}
// PresignGetURL 限时直链下载(委托当前引擎)。
func (m *Manager) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return m.current.PresignGetURL(ctx, savePath, expires)
}
// PresignPutURL 限时直传(委托当前引擎)。
func (m *Manager) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return m.current.PresignPutURL(ctx, savePath, expires)
}
// HealthCheck 健康检查(委托当前引擎)。
func (m *Manager) HealthCheck(ctx context.Context) error {
return m.current.HealthCheck(ctx)
}
// —— 管理面:当前引擎名 / 按名取实例 / 热切换 / 缓存失效 ——
// CurrentName 当前引擎名(管理端展示与文件归属戳用;并发安全)。
func (m *Manager) CurrentName() string {
m.mu.RLock()
defer m.mu.RUnlock()
return m.curName
}
// Current 当前引擎实例。
func (m *Manager) Current() Storage {
m.mu.RLock()
defer m.mu.RUnlock()
return m.current
}
// EngineOf 按名字取引擎实例(带缓存;用于按文件归属引擎取回旧文件)。
// 实例不存在时现场构建(不健康检查——读旧文件尽力而为,构建失败即报错)。
func (m *Manager) EngineOf(name string) (Storage, error) {
if !validEngines[name] {
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav", name)
}
m.mu.RLock()
if s, ok := m.cache[name]; ok {
m.mu.RUnlock()
return s, nil
}
m.mu.RUnlock()
m.mu.Lock()
defer m.mu.Unlock()
// 双检:拿写锁期间可能已被并发构建
if s, ok := m.cache[name]; ok {
return s, nil
}
s, err := m.build(name)
if err != nil {
return nil, err
}
m.cache[name] = s
return s, nil
}
// Switch 热切换当前引擎:构建新实例 → 健康检查 → 成功才替换指针。
// 任一步失败返回错误且当前引擎保持不变(管理端 503 上报)。
func (m *Manager) Switch(name string) (Storage, error) {
if !validEngines[name] {
return nil, fmt.Errorf("storage: 未知存储引擎 %q(仅支持 local|s3|webdav", name)
}
m.mu.Lock()
defer m.mu.Unlock()
if m.curName == name {
return m.current, nil
}
s, ok := m.cache[name]
if !ok {
var err error
s, err = m.build(name)
if err != nil {
return nil, fmt.Errorf("storage: 构建 %s 引擎失败: %w", name, err)
}
}
if err := s.HealthCheck(context.Background()); err != nil {
return nil, fmt.Errorf("storage: %s 引擎健康检查未通过: %w", name, err)
}
m.current = s
m.curName = name
m.cache[name] = s
return s, nil
}
// Invalidate 引擎参数变更后使对应实例缓存失效(下次 EngineOf/Switch 重建生效)。
// 当前引擎不受影响(运行中实例继续服务,直到显式 Switch)。
func (m *Manager) Invalidate(name string) {
m.mu.Lock()
defer m.mu.Unlock()
if name == m.curName {
return // 当前引擎实例仍被热路径使用,不重建;参数生效由下一次 Switch 完成
}
delete(m.cache, name)
}
+197
View File
@@ -0,0 +1,197 @@
package storage
import (
"context"
"errors"
"io"
"strings"
"sync/atomic"
"testing"
)
// fakeEngine 可配置健康检查结果的桩引擎。
type fakeEngine struct{ failHealth bool }
func (f *fakeEngine) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
return 0, nil
}
func (f *fakeEngine) DeleteFile(ctx context.Context, savePath string) error { return nil }
func (f *fakeEngine) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
return nil, ErrNotFound
}
func (f *fakeEngine) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
return nil, ErrNotFound
}
func (f *fakeEngine) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
return 0, nil
}
func (f *fakeEngine) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
return 0, "", nil
}
func (f *fakeEngine) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
return nil
}
func (f *fakeEngine) FileExists(ctx context.Context, savePath string) (bool, error) {
return false, nil
}
func (f *fakeEngine) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
return nil, nil, ErrNotSupported
}
func (f *fakeEngine) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
func (f *fakeEngine) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
func (f *fakeEngine) HealthCheck(ctx context.Context) error {
if f.failHealth {
return ErrUnavailable
}
return nil
}
// newTestManager 构造测试用 Managerlocal 健康引擎起步;s3/webdav 由计数器控制健康。
func newTestManager(s3Fail *atomic.Bool) *Manager {
build := func(name string) (Storage, error) {
switch name {
case "local":
return &fakeEngine{}, nil
case "s3":
return &fakeEngine{failHealth: s3Fail.Load()}, nil
case "webdav":
return &fakeEngine{}, nil
}
return nil, errors.New("unknown")
}
return NewManager("local", &fakeEngine{}, build)
}
// TestSwitchSuccessAndCurrentName 切换成功后当前引擎名与实例更新。
func TestSwitchSuccessAndCurrentName(t *testing.T) {
var s3Fail atomic.Bool
m := newTestManager(&s3Fail)
if m.CurrentName() != "local" {
t.Fatalf("初始引擎应为 local,得到 %s", m.CurrentName())
}
if _, err := m.Switch("s3"); err != nil {
t.Fatalf("Switch(s3) 失败: %v", err)
}
if m.CurrentName() != "s3" {
t.Fatalf("切换后引擎应为 s3,得到 %s", m.CurrentName())
}
if _, err := m.Switch("webdav"); err != nil {
t.Fatalf("Switch(webdav) 失败: %v", err)
}
if m.CurrentName() != "webdav" {
t.Fatalf("切换后引擎应为 webdav,得到 %s", m.CurrentName())
}
}
// TestSwitchFailureKeepsCurrent 健康检查失败时保持原引擎(v3 核心语义)。
func TestSwitchFailureKeepsCurrent(t *testing.T) {
var s3Fail atomic.Bool
s3Fail.Store(true) // s3 不健康
m := newTestManager(&s3Fail)
if _, err := m.Switch("s3"); err == nil {
t.Fatal("s3 不健康时 Switch 应失败")
}
if m.CurrentName() != "local" {
t.Fatalf("切换失败后应保持 local,得到 %s", m.CurrentName())
}
// 恢复健康后可切换成功
s3Fail.Store(false)
if _, err := m.Switch("s3"); err != nil {
t.Fatalf("恢复健康后 Switch(s3) 应成功: %v", err)
}
if m.CurrentName() != "s3" {
t.Fatalf("恢复后引擎应为 s3,得到 %s", m.CurrentName())
}
}
// TestSwitchInvalidName 非法引擎名拒绝。
func TestSwitchInvalidName(t *testing.T) {
m := newTestManager(&atomic.Bool{})
if _, err := m.Switch("ftp"); err == nil || !strings.Contains(err.Error(), "未知存储引擎") {
t.Fatalf("非法引擎名应报未知存储引擎,得到 %v", err)
}
if !ValidEngine("local") || ValidEngine("ftp") {
t.Fatal("ValidEngine 判定错误")
}
}
// TestEngineOfCacheAndInvalidate EngineOf 缓存命中 + Invalidate 后重建(参数生效路径)。
func TestEngineOfCacheAndInvalidate(t *testing.T) {
var builds atomic.Int64
build := func(name string) (Storage, error) {
builds.Add(1)
return &fakeEngine{failHealth: false}, nil
}
m := NewManager("local", &fakeEngine{}, build)
s1, err := m.EngineOf("s3")
if err != nil {
t.Fatalf("EngineOf(s3): %v", err)
}
s2, err := m.EngineOf("s3")
if err != nil {
t.Fatalf("EngineOf(s3) second: %v", err)
}
if s1 != s2 {
t.Fatal("EngineOf 应命中缓存返回同一实例")
}
if n := builds.Load(); n != 1 {
t.Fatalf("应只构建 1 次,实际 %d", n)
}
// Invalidate 后下次取重建新实例
m.Invalidate("s3")
s3, err := m.EngineOf("s3")
if err != nil {
t.Fatalf("EngineOf(s3) after invalidate: %v", err)
}
if s3 == s1 {
t.Fatal("Invalidate 后应返回重建的新实例")
}
if n := builds.Load(); n != 2 {
t.Fatalf("Invalidate 后应再构建 1 次,实际累计 %d", n)
}
}
// TestInvalidateCurrentNoop Invalidate 当前引擎不生效(热路径实例保持)。
func TestInvalidateCurrentNoop(t *testing.T) {
m := newTestManager(&atomic.Bool{})
cur := m.Current()
m.Invalidate("local") // 当前引擎:应为 no-op
if m.Current() != cur {
t.Fatal("Invalidate 当前引擎不应替换实例")
}
}
// TestSwitchSameNameNoop 同名 Switch 幂等。
func TestSwitchSameNameNoop(t *testing.T) {
m := newTestManager(&atomic.Bool{})
s, err := m.Switch("local")
if err != nil {
t.Fatalf("Switch(local) 同名应成功: %v", err)
}
if s != m.Current() {
t.Fatal("同名 Switch 应返回当前实例")
}
}
// TestDelegateToCurrent 保存/读取类操作委托当前引擎(切换后指向新引擎)。
func TestDelegateToCurrent(t *testing.T) {
var s3Fail atomic.Bool
m := newTestManager(&s3Fail)
ctx := context.Background()
// local 引擎 HealthCheck 健康
if err := m.HealthCheck(ctx); err != nil {
t.Fatalf("委托 HealthCheck(local): %v", err)
}
if _, err := m.Switch("s3"); err != nil {
t.Fatalf("Switch(s3): %v", err)
}
if err := m.HealthCheck(ctx); err != nil {
t.Fatalf("委托 HealthCheck(s3): %v", err)
}
}
+70
View File
@@ -0,0 +1,70 @@
package storage
// EngineOptions 引擎构造选项:由 main.goAPI 层任务)从 config KV 填充。
// 各引擎的 RegisterEngine 工厂读取本结构;零值即安全默认。
type EngineOptions struct {
// Local 本地引擎选项。
Local LocalOptions
// S3 S3 引擎选项。
S3 S3Options
// WebDAV WebDAV 引擎选项。
WebDAV WebDAVOptions
}
// LocalOptions 本地引擎配置(对齐 local_storage_path)。
type LocalOptions struct {
// Root 存储根目录;空则使用系统临时目录。
Root string
}
// S3Options S3 引擎配置(对齐 s3_* 配置键)。
type S3Options struct {
AccessKeyID string // s3_access_key_id
SecretAccessKey string // s3_secret_access_key
SessionToken string // aws_session_token
Bucket string // s3_bucket_name
Endpoint string // s3_endpoint_urlMinIO 等;空则 AWS 默认端点)
Region string // s3_region_name,默认 auto
AddressingStyle string // s3_addressing_style: auto|path|virtual
}
// WebDAVOptions WebDAV 引擎配置(对齐 webdav_* 配置键 + 本次优化项)。
type WebDAVOptions struct {
// BaseURL 服务地址,如 https://dav.example.com/dav/。
BaseURL string
// Username/Password 凭据(Basic 与 Digest 共用)。
Username string
Password string
// RootPath 远端根目录(webdav_root_path),会自动逐级创建。
RootPath string
// MaxRetries 5xx/网络错误最大重试次数(指数退避),0 取默认 3。
MaxRetries int
// BaseBackoff 重试基础退避时长,0 取默认 200ms。
BaseBackoff int64
// Timeout 单请求超时秒数,0 取默认 30s。
Timeout int64
// MaxIdleConnsPerHost 连接池每主机最大空闲连接,0 取默认 16(连接复用优化)。
MaxIdleConnsPerHost int
}
// engineOptions 全局引擎选项(由 main.go 注入;默认零值)。
var engineOptions EngineOptions
// SetEngineOptions 注入引擎构造选项(在 RegisterEngine 工厂执行前调用)。
func SetEngineOptions(opts EngineOptions) { engineOptions = opts }
// applyDefaults 填充零值默认项。
func (o *WebDAVOptions) applyDefaults() {
if o.MaxRetries <= 0 {
o.MaxRetries = 3
}
if o.BaseBackoff <= 0 {
o.BaseBackoff = 200
}
if o.Timeout <= 0 {
o.Timeout = 30
}
if o.MaxIdleConnsPerHost <= 0 {
o.MaxIdleConnsPerHost = 16
}
}
+103
View File
@@ -0,0 +1,103 @@
package storage
import (
"path"
"strings"
)
// ChunkDir 实现默认分片目录约定:<父目录>/chunks/<uploadID>。
// local/s3/webdav 三引擎共用,保持分片路径一致。
func ChunkDir(savePath, uploadID string) string {
dir := path.Dir(savePath)
name := path.Base(savePath)
// 防御:savePath 非法时仍返回明确结构,具体引擎再做安全校验
if name == "." || name == "/" {
name = "file"
}
return path.Join(dir, "chunks", uploadID) + "/" + name
}
// ChunkPartPath 分片对象完整路径(相对存储根)。
func ChunkPartPath(savePath, uploadID string, index int) string {
dir := path.Dir(savePath)
return path.Join(dir, "chunks", uploadID, itoa(index)+".part")
}
// SanitizePath 清理相对路径:统一斜杠、去首尾斜杠、拒绝 .. 穿越。
// 返回清理后的相对路径与是否合法。
func SanitizePath(p string) (string, bool) {
raw := strings.ReplaceAll(strings.TrimSpace(p), "\\", "/")
raw = strings.TrimPrefix(raw, "/")
if raw == "" {
return "", false
}
cleaned := path.Clean(raw)
if cleaned == ".." || strings.HasPrefix(cleaned, "../") || path.IsAbs(cleaned) {
return "", false
}
// 拒绝任何单独的 .. 段
for _, seg := range strings.Split(cleaned, "/") {
if seg == ".." {
return "", false
}
}
return cleaned, true
}
// SanitizeFileName 清理文件名:剥离路径、替换非法字符、限制长度。
// 对齐参考 core/utils.py 的 sanitize_filename。
func SanitizeFileName(name string) string {
// 剥离路径
if idx := strings.LastIndexAny(name, "/\\"); idx >= 0 {
name = name[idx+1:]
}
var b strings.Builder
for _, r := range name {
switch {
case r < 0x20 || r == 0x7f:
b.WriteByte('_')
case strings.ContainsRune(`\*?:"<>|`, r):
b.WriteByte('_')
case r == ' ':
b.WriteByte('_')
default:
b.WriteRune(r)
}
}
cleaned := b.String()
// 压缩连续下划线
for strings.Contains(cleaned, "__") {
cleaned = strings.ReplaceAll(cleaned, "__", "_")
}
cleaned = strings.Trim(cleaned, "._")
if cleaned == "" {
return "unnamed_file"
}
if len(cleaned) > 255 {
cleaned = cleaned[:255]
}
return cleaned
}
// itoa 小整数转字符串。
func itoa(n int) 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:])
}
+74
View File
@@ -0,0 +1,74 @@
package storage
import (
"crypto/sha256"
"encoding/hex"
"io"
"testing"
)
// TestSanitizePath 校验路径穿越防护。
func TestSanitizePath(t *testing.T) {
cases := []struct {
in string
ok bool
out string
}{
{"2025/08/uuid.zip", true, "2025/08/uuid.zip"},
{"/2025/08/uuid.zip", true, "2025/08/uuid.zip"},
{"a\\b\\c.txt", true, "a/b/c.txt"},
{"../etc/passwd", false, ""},
{"a/../../b", false, ""},
{"..", false, ""},
{"", false, ""},
}
for _, tc := range cases {
got, ok := SanitizePath(tc.in)
if ok != tc.ok || (ok && got != tc.out) {
t.Errorf("SanitizePath(%q) = (%q, %v), want (%q, %v)", tc.in, got, ok, tc.out, tc.ok)
}
}
}
// TestSanitizeFileName 校验文件名清理。
func TestSanitizeFileName(t *testing.T) {
cases := []struct{ in, want string }{
{"hello world.zip", "hello_world.zip"},
{"/path/to/file.txt", "file.txt"},
{"a<b>:c?.mp4", "a_b_c_.mp4"}, // 连续下划线压缩,对齐参考 re.sub(r"_+", "_")
{"", "unnamed_file"},
{"__..__", "unnamed_file"},
}
for _, tc := range cases {
if got := SanitizeFileName(tc.in); got != tc.want {
t.Errorf("SanitizeFileName(%q) = %q, want %q", tc.in, got, tc.want)
}
}
}
// TestChunkPartPath 校验分片路径约定。
func TestChunkPartPath(t *testing.T) {
got := ChunkPartPath("2025/08/uuid.zip", "upload-1", 3)
want := "2025/08/chunks/upload-1/3.part"
if got != want {
t.Errorf("ChunkPartPath = %q, want %q", got, want)
}
}
// TestChunkDir 校验分片目录约定。
func TestChunkDir(t *testing.T) {
got := ChunkDir("2025/08/uuid.zip", "upload-1")
want := "2025/08/chunks/upload-1/uuid.zip"
if got != want {
t.Errorf("ChunkDir = %q, want %q", got, want)
}
}
// TestSHA256Helper 辅助:确认 sha256 用法一致(合并校验依赖)。
func TestSHA256Helper(t *testing.T) {
h := sha256.Sum256([]byte("abc"))
if got := hex.EncodeToString(h[:]); got != "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad" {
t.Errorf("sha256(abc) = %s", got)
}
_ = io.EOF
}
+650
View File
@@ -0,0 +1,650 @@
package storage
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"net"
"net/http"
"os"
"strconv"
"strings"
"time"
"github.com/aws/aws-sdk-go-v2/aws"
awshttp "github.com/aws/aws-sdk-go-v2/aws/transport/http"
"github.com/aws/aws-sdk-go-v2/config"
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/feature/s3/manager"
"github.com/aws/aws-sdk-go-v2/service/s3"
"github.com/aws/aws-sdk-go-v2/service/s3/types"
"github.com/aws/smithy-go"
)
// S3Storage 基于 aws-sdk-go-v2 的 S3 兼容对象存储引擎(AWS / MinIO / R2 / OSS 等)。
//
// 相比参考实现(S3FileStorageaioboto3)的改进:
// - 单例客户端 + 自定义连接池 Transport(参考实现每次操作新建 session);
// - SaveFile 走 manager.Uploader 分片并发上传(未知长度也可流式,内存占用 ≤ partSize);
// - SaveChunk 落本地临时文件获取精确 Content-Length(参考实现整块读入内存);
// - MergeChunks 用 S3 原生 multipart 流式合并,边读边校验哈希(不落盘、不整块进内存);
// - 5xx/网络错误由 SDK 内置指数退避重试器处理(可配次数)。
type S3Storage struct {
client *s3.Client
presigner *s3.PresignClient
uploader *manager.Uploader
bucket string
}
// NewS3Storage 构造 S3 引擎。
func NewS3Storage(opts S3Options) (*S3Storage, error) {
if strings.TrimSpace(opts.Bucket) == "" {
return nil, fmt.Errorf("storage/s3: 缺少 bucket 配置(s3_bucket_name")
}
region := strings.TrimSpace(opts.Region)
if region == "" {
region = "us-east-1"
}
loadOpts := []func(*config.LoadOptions) error{
config.WithRegion(region),
config.WithCredentialsProvider(credentials.NewStaticCredentialsProvider(
opts.AccessKeyID, opts.SecretAccessKey, opts.SessionToken,
)),
// SDK 内置重试器:标准模式,指数退避 + 抖动,覆盖 5xx 与网络错误。
config.WithRetryMaxAttempts(3),
// 兼容性:仅协议要求时才计算校验和。默认的 trailing CRC32 需要可重放流
// 或 TLSMinIO/R2 等自建端点通常不需要,关闭后 MergeChunks 的
// GET→UploadPart 纯流式转发才能工作。
config.WithRequestChecksumCalculation(aws.RequestChecksumCalculationWhenRequired),
config.WithResponseChecksumValidation(aws.ResponseChecksumValidationWhenRequired),
}
if ep := strings.TrimSpace(opts.Endpoint); ep != "" {
loadOpts = append(loadOpts, config.WithBaseEndpoint(ep))
}
awsCfg, err := config.LoadDefaultConfig(context.Background(), loadOpts...)
if err != nil {
return nil, fmt.Errorf("storage/s3: 初始化 SDK 配置失败: %w", err)
}
client := s3.NewFromConfig(awsCfg, func(o *s3.Options) {
// 寻址风格:path 显式启用;auto 时自定义端点(自建 MinIO 等)默认 path-style。
switch strings.ToLower(strings.TrimSpace(opts.AddressingStyle)) {
case "path":
o.UsePathStyle = true
case "virtual":
o.UsePathStyle = false
default: // auto
o.UsePathStyle = strings.TrimSpace(opts.Endpoint) != ""
}
// 连接复用:自定义 Transport 连接池。
o.HTTPClient = newPooledHTTPClient()
})
st := &S3Storage{
client: client,
presigner: s3.NewPresignClient(client),
bucket: opts.Bucket,
}
st.uploader = manager.NewUploader(client, func(u *manager.Uploader) {
u.PartSize = 5 * 1024 * 1024 // 5MBS3 multipart 最小分片
u.Concurrency = 4
u.LeavePartsOnError = false
})
return st, nil
}
func init() {
RegisterEngine("s3", func(ctx context.Context) (Storage, error) {
return NewS3Storage(engineOptions.S3)
})
}
// newPooledHTTPClient 供 SDK 使用的连接池化 HTTP 客户端。
func newPooledHTTPClient() *http.Client {
return &http.Client{
Transport: &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: 10 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
MaxIdleConnsPerHost: 16,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: time.Second,
ResponseHeaderTimeout: 60 * time.Second,
},
}
}
// key 校验并规范化对象键(拒绝穿越,统一斜杠)。
func (s *S3Storage) key(savePath string) (string, error) {
cleaned, ok := SanitizePath(savePath)
if !ok {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
return cleaned, nil
}
// SaveFile 流式保存:manager.Uploader 按需分片并发上传,内存占用恒定。
func (s *S3Storage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
key, err := s.key(savePath)
if err != nil {
return 0, err
}
src := &countingReader{r: r}
_, err = s.uploader.Upload(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
Body: src,
ContentType: aws.String("application/octet-stream"),
})
if err != nil {
return src.count(), mapS3Error(err, "PutObject")
}
return src.count(), nil
}
// DeleteFile 删除对象;S3 对不存在的键也返回成功。
func (s *S3Storage) DeleteFile(ctx context.Context, savePath string) error {
key, err := s.key(savePath)
if err != nil {
return err
}
_, err = s.client.DeleteObject(ctx, &s3.DeleteObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
})
return mapS3Error(err, "DeleteObject")
}
// Open 获取下载流:Range 直接透传为 GetObject Range 头(流式,不落盘)。
func (s *S3Storage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
key, err := s.key(savePath)
if err != nil {
return nil, err
}
input := &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
}
if rng != nil {
input.Range = aws.String(rangeHeaderValue(rng))
}
out, err := s.client.GetObject(ctx, input)
if err != nil {
return nil, mapS3Error(err, "GetObject")
}
size := aws.ToInt64(out.ContentLength)
start, end := int64(0), size-1
if cr := aws.ToString(out.ContentRange); cr != "" { // 服务端按 206 返回了区间
if sr, e, total, ok := parseContentRange(cr); ok {
start, end = sr, e
if total >= 0 {
size = total
}
}
}
return &Download{
ReadCloser: out.Body,
Start: start,
End: end,
Total: size,
Meta: FileMeta{
Size: size,
ContentType: aws.ToString(out.ContentType),
AcceptRanges: true,
},
}, nil
}
// rangeHeaderValue 将 Range 结构转为 HTTP Range 头值。
func rangeHeaderValue(rng *Range) string {
if rng.End < 0 {
return fmt.Sprintf("bytes=%d-", rng.Start)
}
return fmt.Sprintf("bytes=%d-%d", rng.Start, rng.End)
}
// parseContentRange 解析 "bytes 0-99/1000"total 可能为 "*")。
func parseContentRange(v string) (start, end, total int64, ok bool) {
v = strings.TrimSpace(v)
if !strings.HasPrefix(v, "bytes ") {
return 0, 0, -1, false
}
parts := strings.SplitN(strings.TrimPrefix(v, "bytes "), "/", 2)
if len(parts) != 2 {
return 0, 0, -1, false
}
total = -1
if parts[1] != "*" {
t, err := strconv.ParseInt(parts[1], 10, 64)
if err != nil {
return 0, 0, -1, false
}
total = t
}
se := strings.SplitN(parts[0], "-", 2)
if len(se) != 2 {
return 0, 0, -1, false
}
s0, err1 := strconv.ParseInt(se[0], 10, 64)
e0, err2 := strconv.ParseInt(se[1], 10, 64)
if err1 != nil || err2 != nil {
return 0, 0, -1, false
}
return s0, e0, total, true
}
// Stat 获取对象元信息(HeadObject)。
func (s *S3Storage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
key, err := s.key(savePath)
if err != nil {
return nil, err
}
out, err := s.client.HeadObject(ctx, &s3.HeadObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
})
if err != nil {
return nil, mapS3Error(err, "HeadObject")
}
return &FileMeta{
Size: aws.ToInt64(out.ContentLength),
ContentType: aws.ToString(out.ContentType),
AcceptRanges: true,
}, nil
}
// SaveChunk 保存分片对象:落临时文件获取精确长度后 PutObject(可重试)。
func (s *S3Storage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
if _, err := s.key(savePath); err != nil {
return 0, err
}
key, err := s.key(ChunkPartPath(savePath, uploadID, chunkIndex))
if err != nil {
return 0, err
}
// 落临时文件:获得精确 Content-Length 与可重放 Body(网络失败可安全重试)。
tmp, err := os.CreateTemp("", "fcb-s3-chunk-*")
if err != nil {
return 0, fmt.Errorf("storage/s3: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
size, err := io.Copy(tmp, r)
if err != nil {
return 0, fmt.Errorf("storage/s3: 缓存分片失败: %w", err)
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, fmt.Errorf("storage/s3: 回卷分片失败: %w", err)
}
_, err = s.client.PutObject(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
Body: tmp,
ContentLength: aws.Int64(size),
ContentType: aws.String("application/octet-stream"),
})
if err != nil {
return size, mapS3Error(err, "PutObject(分片)")
}
return size, nil
}
// S3 multipart 最小分片限制:除最后一片外每片 ≥5MB,否则 Complete 返回 EntityTooSmall。
// 分片上传的分块大小由服务端配置保证(建议 ≥5MB)。
const s3MinPartSize = 5 * 1024 * 1024
// MergeChunks 用 S3 原生 multipart 流式合并:
// 逐分片 GET → 边流边算哈希 → UploadPart(带精确 Content-Length)→ Complete。
// 任一步失败即 Abort 并返回错误;成功后清理分片对象。
func (s *S3Storage) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
if total <= 0 {
return 0, "", fmt.Errorf("storage/s3: 非法分片总数 %d", total)
}
key, err := s.key(savePath)
if err != nil {
return 0, "", err
}
chunkPrefix, err := s.key(chunkDirOf(savePath, uploadID))
if err != nil {
return 0, "", err
}
// 单分片快速路径:直接流式 PutObject,绕过 multipart 的 5MB 限制。
if total == 1 {
return s.mergeSingle(ctx, chunkPrefix+"/0.part", key, verifyHash, 0)
}
mpu, err := s.client.CreateMultipartUpload(ctx, &s3.CreateMultipartUploadInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
ContentType: aws.String("application/octet-stream"),
})
if err != nil {
return 0, "", mapS3Error(err, "CreateMultipartUpload")
}
_ = aws.ToString(mpu.UploadId) // S3 侧 multipart 会话 IDAbort 时复用 mpu.UploadId
size := int64(0)
totalHash := sha256.New()
parts := make([]types.CompletedPart, 0, total)
defer func() {
// 出错时取消 multipart(避免残留分片产生存储费用)。
if len(parts) < total {
_, _ = s.client.AbortMultipartUpload(ctx, &s3.AbortMultipartUploadInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
UploadId: mpu.UploadId,
})
}
}()
for i := 0; i < total; i++ {
if err := ctx.Err(); err != nil {
return 0, "", err
}
var expected string
if verifyHash != nil {
expected, err = verifyHash(i)
if err != nil {
return 0, "", err
}
}
getOut, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(fmt.Sprintf("%s/%d.part", chunkPrefix, i)),
})
if err != nil {
return 0, "", mapS3Error(err, fmt.Sprintf("GetObject(分片 %d)", i))
}
// 分片先流式落临时文件:计算哈希 + 获得可回卷 body(SDK 签名哈希需要 seekable 流,
// 同时为 UploadPart 失败重试保留数据)。
chunkHash := sha256.New()
tmp, err := os.CreateTemp("", "fcb-s3-part-*")
if err != nil {
_ = getOut.Body.Close()
return 0, "", fmt.Errorf("storage/s3: 创建分片临时文件失败: %w", err)
}
partLen, err := io.Copy(io.MultiWriter(tmp, totalHash, chunkHash), getOut.Body)
_ = getOut.Body.Close()
if err != nil {
_ = tmp.Close()
_ = os.Remove(tmp.Name())
return 0, "", fmt.Errorf("storage/s3: 读取分片 %d 失败: %w", i, err)
}
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
_ = tmp.Close()
_ = os.Remove(tmp.Name())
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
_ = tmp.Close()
_ = os.Remove(tmp.Name())
return 0, "", fmt.Errorf("storage/s3: 回卷分片 %d 失败: %w", i, err)
}
up, err := s.client.UploadPart(ctx, &s3.UploadPartInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
UploadId: mpu.UploadId,
PartNumber: aws.Int32(int32(i + 1)),
Body: tmp,
ContentLength: aws.Int64(partLen),
})
_ = tmp.Close()
_ = os.Remove(tmp.Name())
if err != nil {
return 0, "", mapS3Error(err, fmt.Sprintf("UploadPart(分片 %d)", i))
}
parts = append(parts, types.CompletedPart{
PartNumber: aws.Int32(int32(i + 1)),
ETag: up.ETag,
})
size += partLen
}
if _, err := s.client.CompleteMultipartUpload(ctx, &s3.CompleteMultipartUploadInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
UploadId: mpu.UploadId,
MultipartUpload: &types.CompletedMultipartUpload{Parts: parts},
}); err != nil {
return 0, "", mapS3Error(err, "CompleteMultipartUpload")
}
// 合并成功后清理分片对象(静默容错)。
_ = s.CleanChunks(ctx, uploadID, savePath)
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
}
// mergeSingle 单分片合并快速路径:GET 分片 → 落临时文件校验 → PutObject 正式键。
func (s *S3Storage) mergeSingle(ctx context.Context, chunkKey, dstKey string, verifyHash func(index int) (string, error), index int) (int64, string, error) {
getOut, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(chunkKey),
})
if err != nil {
return 0, "", mapS3Error(err, "GetObject(分片)")
}
defer func() { _ = getOut.Body.Close() }()
tmp, err := os.CreateTemp("", "fcb-s3-merge-*")
if err != nil {
return 0, "", fmt.Errorf("storage/s3: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
fileHash := sha256.New()
size, err := io.Copy(io.MultiWriter(tmp, fileHash), getOut.Body)
if err != nil {
return 0, "", fmt.Errorf("storage/s3: 读取分片失败: %w", err)
}
if verifyHash != nil {
expected, err := verifyHash(index)
if err != nil {
return 0, "", err
}
if expected != "" && expected != hex.EncodeToString(fileHash.Sum(nil)) {
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, index, expected)
}
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, "", fmt.Errorf("storage/s3: 回卷临时文件失败: %w", err)
}
if _, err := s.client.PutObject(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(dstKey),
Body: tmp,
ContentLength: aws.Int64(size),
ContentType: aws.String("application/octet-stream"),
}); err != nil {
return 0, "", mapS3Error(err, "PutObject(合并)")
}
return size, hex.EncodeToString(fileHash.Sum(nil)), nil
}
// CleanChunks 列举并批量删除分片对象;前缀不存在时静默成功。
func (s *S3Storage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
prefix, err := s.key(chunkDirOf(savePath, uploadID))
if err != nil {
return err
}
prefix += "/"
paginator := s3.NewListObjectsV2Paginator(s.client, &s3.ListObjectsV2Input{
Bucket: aws.String(s.bucket),
Prefix: aws.String(prefix),
})
for paginator.HasMorePages() {
page, err := paginator.NextPage(ctx)
if err != nil {
return mapS3Error(err, "ListObjectsV2(分片)")
}
if len(page.Contents) == 0 {
return nil
}
objs := make([]types.ObjectIdentifier, 0, len(page.Contents))
for _, obj := range page.Contents {
objs = append(objs, types.ObjectIdentifier{Key: obj.Key})
}
if _, err := s.client.DeleteObjects(ctx, &s3.DeleteObjectsInput{
Bucket: aws.String(s.bucket),
Delete: &types.Delete{Objects: objs, Quiet: aws.Bool(true)},
}); err != nil {
return mapS3Error(err, "DeleteObjects(分片)")
}
}
return nil
}
// FileExists HeadObject 探测存在性。
func (s *S3Storage) FileExists(ctx context.Context, savePath string) (bool, error) {
key, err := s.key(savePath)
if err != nil {
return false, err
}
_, err = s.client.HeadObject(ctx, &s3.HeadObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
})
if err != nil {
if isS3NotFound(err) {
return false, nil
}
return false, mapS3Error(err, "HeadObject")
}
return true, nil
}
// PresignGetURL 生成限时下载直链。
func (s *S3Storage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
key, err := s.key(savePath)
if err != nil {
return "", err
}
if expires <= 0 {
expires = 3600
}
out, err := s.presigner.PresignGetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
}, s3.WithPresignExpires(time.Duration(expires)*time.Second))
if err != nil {
return "", mapS3Error(err, "PresignGetObject")
}
return out.URL, nil
}
// PresignPutURL 生成限时直传 URL。
func (s *S3Storage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
key, err := s.key(savePath)
if err != nil {
return "", err
}
if expires <= 0 {
expires = 900
}
out, err := s.presigner.PresignPutObject(ctx, &s3.PutObjectInput{
Bucket: aws.String(s.bucket),
Key: aws.String(key),
}, s3.WithPresignExpires(time.Duration(expires)*time.Second))
if err != nil {
return "", mapS3Error(err, "PresignPutObject")
}
return out.URL, nil
}
// HealthCheck 健康检查:列举 bucketMaxKeys=1),同时校验连通性、凭据与 bucket 存在。
func (s *S3Storage) HealthCheck(ctx context.Context) error {
maxKeys := int32(1)
_, err := s.client.ListObjectsV2(ctx, &s3.ListObjectsV2Input{
Bucket: aws.String(s.bucket),
MaxKeys: aws.Int32(maxKeys),
Prefix: aws.String(""),
})
if err != nil {
return fmt.Errorf("%w: S3 健康检查失败: %v", ErrUnavailable, err)
}
return nil
}
// ReadHead 读取对象前 n 字节(保留的便捷封装:HeadMeta 的仅头部形态)。
func (s *S3Storage) ReadHead(ctx context.Context, savePath string, n int64) ([]byte, error) {
_, head, err := s.HeadMeta(ctx, savePath, n)
return head, err
}
// HeadMeta 读取对象元信息与头部字节(S3 引擎实现:HeadObject + Range GET)。
func (s *S3Storage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
key, err := s.key(savePath)
if err != nil {
return nil, nil, err
}
head, err := s.headBytes(ctx, s.bucket, key, headBytes)
if err != nil {
return nil, nil, err
}
meta, err := s.Stat(ctx, savePath)
if err != nil {
return nil, nil, err
}
return meta, head, nil
}
// headBytes 通过 Range GET 读取对象前 n 字节。
func (s *S3Storage) headBytes(ctx context.Context, bucket, key string, n int64) ([]byte, error) {
if n <= 0 {
return nil, nil
}
out, err := s.client.GetObject(ctx, &s3.GetObjectInput{
Bucket: aws.String(bucket),
Key: aws.String(key),
Range: aws.String(fmt.Sprintf("bytes=0-%d", n-1)),
})
if err != nil {
return nil, mapS3Error(err, "GetObject(head)")
}
defer func() { _ = out.Body.Close() }()
return io.ReadAll(io.LimitReader(out.Body, n))
}
// isS3NotFound 判断错误是否为对象不存在。
func isS3NotFound(err error) bool {
var nf *types.NotFound
if errors.As(err, &nf) {
return true
}
var ae smithy.APIError
if errors.As(err, &ae) {
switch ae.ErrorCode() {
case "NotFound", "NoSuchKey":
return true
}
}
var re *awshttp.ResponseError
if errors.As(err, &re) && re.HTTPStatusCode() == http.StatusNotFound {
return true
}
return false
}
// mapS3Error 将 SDK 错误映射为包内哨兵错误。
func mapS3Error(err error, op string) error {
if err == nil {
return nil
}
if isS3NotFound(err) {
return fmt.Errorf("%w%s", ErrNotFound, op)
}
var re *awshttp.ResponseError
if errors.As(err, &re) && re.HTTPStatusCode() == http.StatusRequestedRangeNotSatisfiable {
return fmt.Errorf("%w%s", ErrRangeNotSatisfiable, op)
}
var ae smithy.APIError
if errors.As(err, &ae) && ae.ErrorCode() == "InvalidRange" {
return fmt.Errorf("%w%s", ErrRangeNotSatisfiable, op)
}
return fmt.Errorf("storage/s3: %s 失败: %w", op, err)
}
// 接口编译期断言。
var _ Storage = (*S3Storage)(nil)
+495
View File
@@ -0,0 +1,495 @@
package storage
import (
"bytes"
"context"
"encoding/xml"
"fmt"
"io"
"net/http"
"net/http/httptest"
"sort"
"strings"
"sync"
"testing"
"time"
)
// ---- 最小 S3 兼容假服务(仅覆盖本引擎用到的 API)----
type fakeS3Upload struct {
key string
parts map[int][]byte
}
type fakeS3 struct {
mu sync.Mutex
objects map[string][]byte
uploads map[string]*fakeS3Upload
nextID int
putCount int
getCount int
headCount int
deleteCount int
listCount int
completeN int
// failNextGet:让接下来 N 次 GET 返回 503(重试测试用)。
failNextGet int
}
func newFakeS3() *fakeS3 {
return &fakeS3{objects: map[string][]byte{}, uploads: map[string]*fakeS3Upload{}}
}
// s3Key 从 path-style 路径剥离 bucket 前缀得到对象键。
func s3Key(r *http.Request) (bucket, key string) {
p := strings.TrimPrefix(r.URL.Path, "/")
if i := strings.Index(p, "/"); i >= 0 {
return p[:i], p[i+1:]
}
return p, ""
}
func s3ErrorXML(w http.ResponseWriter, status int, code, msg string) {
w.Header().Set("Content-Type", "application/xml")
w.WriteHeader(status)
_, _ = w.Write([]byte(fmt.Sprintf(
`<?xml version="1.0" encoding="UTF-8"?><Error><Code>%s</Code><Message>%s</Message></Error>`, code, msg)))
}
func (f *fakeS3) ServeHTTP(w http.ResponseWriter, r *http.Request) {
f.mu.Lock()
defer f.mu.Unlock()
q := r.URL.Query()
bucket, key := s3Key(r)
switch {
// UploadPart
case r.Method == http.MethodPut && q.Get("partNumber") != "" && q.Get("uploadId") != "":
up, ok := f.uploads[q.Get("uploadId")]
if !ok {
s3ErrorXML(w, 404, "NoSuchUpload", "upload not found")
return
}
body, _ := io.ReadAll(r.Body)
var n int
_, _ = fmt.Sscanf(q.Get("partNumber"), "%d", &n)
up.parts[n] = body
w.Header().Set("ETag", fmt.Sprintf(`"part-%d"`, n))
w.WriteHeader(http.StatusOK)
// CreateMultipartUpload
case r.Method == http.MethodPost && q.Has("uploads"):
f.nextID++
id := fmt.Sprintf("mpu-%d", f.nextID)
f.uploads[id] = &fakeS3Upload{key: key, parts: map[int][]byte{}}
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(fmt.Sprintf(
`<?xml version="1.0" encoding="UTF-8"?><InitiateMultipartUploadResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Bucket>%s</Bucket><Key>%s</Key><UploadId>%s</UploadId></InitiateMultipartUploadResult>`,
bucket, key, id)))
// CompleteMultipartUpload
case r.Method == http.MethodPost && q.Get("uploadId") != "":
up, ok := f.uploads[q.Get("uploadId")]
if !ok {
s3ErrorXML(w, 404, "NoSuchUpload", "upload not found")
return
}
// 按 partNumber 有序拼接
nums := make([]int, 0, len(up.parts))
for n := range up.parts {
nums = append(nums, n)
}
sort.Ints(nums)
var merged bytes.Buffer
for _, n := range nums {
merged.Write(up.parts[n])
}
f.objects[up.key] = merged.Bytes()
delete(f.uploads, q.Get("uploadId"))
f.completeN++
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(fmt.Sprintf(
`<?xml version="1.0" encoding="UTF-8"?><CompleteMultipartUploadResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Location>http://%s/%s/%s</Location><Bucket>%s</Bucket><Key>%s</Key><ETag>"merged"</ETag></CompleteMultipartUploadResult>`,
r.Host, bucket, up.key, bucket, up.key)))
// AbortMultipartUpload
case r.Method == http.MethodDelete && q.Get("uploadId") != "":
delete(f.uploads, q.Get("uploadId"))
w.WriteHeader(http.StatusNoContent)
// DeleteObjects(批量)
case r.Method == http.MethodPost && q.Has("delete"):
var req struct {
Objects []struct {
Key string `xml:"Key"`
} `xml:"Object"`
}
_ = xml.NewDecoder(r.Body).Decode(&req)
for _, o := range req.Objects {
delete(f.objects, o.Key)
}
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(`<?xml version="1.0" encoding="UTF-8"?><DeleteResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"></DeleteResult>`))
// ListObjectsV2
case r.Method == http.MethodGet && q.Get("list-type") == "2":
f.listCount++
prefix := q.Get("prefix")
var body strings.Builder
body.WriteString(`<?xml version="1.0" encoding="UTF-8"?><ListBucketResult xmlns="http://s3.amazonaws.com/doc/2006-03-01/"><Name>` + bucket + `</Name><Prefix>` + prefix + `</Prefix><IsTruncated>false</IsTruncated>`)
keys := make([]string, 0, len(f.objects))
for k := range f.objects {
if strings.HasPrefix(k, prefix) {
keys = append(keys, k)
}
}
sort.Strings(keys)
for _, k := range keys {
body.WriteString(fmt.Sprintf("<Contents><Key>%s</Key><Size>%d</Size></Contents>", k, len(f.objects[k])))
}
body.WriteString("</ListBucketResult>")
w.Header().Set("Content-Type", "application/xml")
_, _ = w.Write([]byte(body.String()))
// PutObject
case r.Method == http.MethodPut:
f.putCount++
body, _ := io.ReadAll(r.Body)
f.objects[key] = body
w.Header().Set("ETag", `"put"`)
w.WriteHeader(http.StatusOK)
// GetObject
case r.Method == http.MethodGet:
if f.failNextGet > 0 {
f.failNextGet--
s3ErrorXML(w, 503, "ServiceUnavailable", "flaky")
return
}
f.getCount++
data, ok := f.objects[key]
if !ok {
s3ErrorXML(w, 404, "NoSuchKey", "not found")
return
}
w.Header().Set("Content-Type", "application/octet-stream")
if rng := r.Header.Get("Range"); rng != "" {
start, end := int64(0), int64(len(data))-1
if _, err := fmt.Sscanf(rng, "bytes=%d-%d", &start, &end); err != nil {
var s int64
if _, err := fmt.Sscanf(rng, "bytes=%d-", &s); err == nil {
start, end = s, int64(len(data))-1
}
}
if start < 0 || start >= int64(len(data)) {
s3ErrorXML(w, 416, "InvalidRange", "range not satisfiable")
return
}
if end >= int64(len(data)) {
end = int64(len(data)) - 1
}
slice := data[start : end+1]
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(data)))
w.WriteHeader(http.StatusPartialContent)
_, _ = w.Write(slice)
return
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write(data)
// HeadObject
case r.Method == http.MethodHead:
f.headCount++
data, ok := f.objects[key]
if !ok {
w.WriteHeader(http.StatusNotFound)
return
}
w.Header().Set("Content-Type", "application/octet-stream")
w.Header().Set("Content-Length", fmt.Sprintf("%d", len(data)))
w.WriteHeader(http.StatusOK)
// DeleteObject
case r.Method == http.MethodDelete:
f.deleteCount++
delete(f.objects, key)
w.WriteHeader(http.StatusNoContent)
default:
s3ErrorXML(w, 400, "NotImplemented", "unsupported")
}
}
// newTestS3 构造对接假服务的 S3 引擎。
func newTestS3(t *testing.T) (*S3Storage, *fakeS3) {
t.Helper()
f := newFakeS3()
srv := httptest.NewServer(f)
t.Cleanup(srv.Close)
st, err := NewS3Storage(S3Options{
AccessKeyID: "test-ak",
SecretAccessKey: "test-sk",
Bucket: "test-bucket",
Endpoint: srv.URL,
Region: "us-east-1",
AddressingStyle: "path",
})
if err != nil {
t.Fatalf("NewS3Storage: %v", err)
}
return st, f
}
// TestS3SaveStatOpenRange 保存/元信息/完整与 Range 下载。
func TestS3SaveStatOpenRange(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
data := []byte("0123456789abcdef S3 引擎测试数据")
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/s3.bin")
if err != nil {
t.Fatalf("SaveFile: %v", err)
}
if n != int64(len(data)) {
t.Fatalf("n = %d", n)
}
if got := f.objects["2025/08/s3.bin"]; !bytes.Equal(got, data) {
t.Fatalf("stored mismatch")
}
meta, err := st.Stat(ctx, "2025/08/s3.bin")
if err != nil {
t.Fatalf("Stat: %v", err)
}
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
t.Fatalf("Stat = %+v", meta)
}
dl, err := st.Open(ctx, "2025/08/s3.bin", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("full content mismatch")
}
if dl.Start != 0 || dl.End != int64(len(data))-1 || dl.Total != int64(len(data)) {
t.Fatalf("full offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
dl, err = st.Open(ctx, "2025/08/s3.bin", &Range{Start: 4, End: 9})
if err != nil {
t.Fatalf("Open range: %v", err)
}
got, _ = io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data[4:10]) {
t.Fatalf("range mismatch: %q", got)
}
if dl.Start != 4 || dl.End != 9 || dl.Total != int64(len(data)) {
t.Fatalf("range offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// 404 / 416
if _, err := st.Open(ctx, "missing.bin", nil); err == nil || !strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("want ErrNotFound, got %v", err)
}
if _, err := st.Open(ctx, "2025/08/s3.bin", &Range{Start: int64(len(data)) + 5, End: -1}); err == nil ||
!strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
}
}
// TestS3DeleteExists 删除与存在性。
func TestS3DeleteExists(t *testing.T) {
st, _ := newTestS3(t)
ctx := context.Background()
if _, err := st.SaveFile(ctx, strings.NewReader("x"), "del.bin"); err != nil {
t.Fatal(err)
}
ok, err := st.FileExists(ctx, "del.bin")
if err != nil || !ok {
t.Fatalf("exists = %v %v", ok, err)
}
if err := st.DeleteFile(ctx, "del.bin"); err != nil {
t.Fatalf("DeleteFile: %v", err)
}
ok, err = st.FileExists(ctx, "del.bin")
if err != nil || ok {
t.Fatalf("after delete exists = %v %v", ok, err)
}
}
// TestS3ChunkMergeMulti 多分片合并:原生 multipart + 哈希校验 + 分片清理。
func TestS3ChunkMergeMulti(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
savePath := "2025/09/s3-chunked.bin"
uploadID := "uid-s3"
chunks := [][]byte{bytes.Repeat([]byte("A"), 6*1024*1024/3), []byte("BBBB"), []byte("CC")}
// 注意:multipart 除最后一片需 ≥5MB;此处只验证代码路径,真实约束由部署配置保证。
// 为避免 EntityTooSmall,将第一片放大:
chunks[0] = bytes.Repeat([]byte("A"), 5*1024*1024)
hashes := make([]string, len(chunks))
var total int64
for i, c := range chunks {
n, err := st.SaveChunk(ctx, uploadID, i, bytes.NewReader(c), savePath)
if err != nil {
t.Fatalf("SaveChunk %d: %v", i, err)
}
hashes[i] = sha256Hex(c)
total += int64(len(c))
_ = n
}
size, fileHash, err := st.MergeChunks(ctx, uploadID, len(chunks), func(i int) (string, error) {
return hashes[i], nil
}, savePath)
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != total {
t.Fatalf("size = %d want %d", size, total)
}
if fileHash != sha256Hex(bytes.Join(chunks, nil)) {
t.Fatalf("file hash mismatch")
}
merged := f.objects[savePath]
if !bytes.Equal(merged, bytes.Join(chunks, nil)) {
t.Fatalf("merged object mismatch (len=%d)", len(merged))
}
if f.completeN != 1 {
t.Fatalf("CompleteMultipartUpload 次数 = %d", f.completeN)
}
// 分片对象已清理
for k := range f.objects {
if strings.Contains(k, "chunks/"+uploadID) {
t.Fatalf("分片对象残留: %s", k)
}
}
}
// TestS3ChunkMergeSingle 单分片快速路径。
func TestS3ChunkMergeSingle(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
data := []byte("single-chunk")
if _, err := st.SaveChunk(ctx, "uid1", 0, bytes.NewReader(data), "one.bin"); err != nil {
t.Fatal(err)
}
size, fileHash, err := st.MergeChunks(ctx, "uid1", 1, func(i int) (string, error) {
return sha256Hex(data), nil
}, "one.bin")
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != int64(len(data)) || fileHash != sha256Hex(data) {
t.Fatalf("size/hash mismatch")
}
if !bytes.Equal(f.objects["one.bin"], data) {
t.Fatalf("object mismatch")
}
}
// TestS3CleanChunks 清理残留分片。
func TestS3CleanChunks(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
savePath := "clean.bin"
for i := 0; i < 3; i++ {
if _, err := st.SaveChunk(ctx, "uidc", i, strings.NewReader("zz"), savePath); err != nil {
t.Fatal(err)
}
}
prefix := chunkDirOf(savePath, "uidc") + "/"
count := 0
for k := range f.objects {
if strings.HasPrefix(k, prefix) {
count++
}
}
if count != 3 {
t.Fatalf("期望 3 个分片对象,实际 %d", count)
}
if err := st.CleanChunks(ctx, "uidc", savePath); err != nil {
t.Fatalf("CleanChunks: %v", err)
}
for k := range f.objects {
if strings.HasPrefix(k, prefix) {
t.Fatalf("分片未清理: %s", k)
}
}
// 幂等:再清理一次不报错
if err := st.CleanChunks(ctx, "uidc", savePath); err != nil {
t.Fatalf("CleanChunks idempotent: %v", err)
}
}
// TestS3Presign 预签名 URL 生成。
func TestS3Presign(t *testing.T) {
st, _ := newTestS3(t)
ctx := context.Background()
getURL, err := st.PresignGetURL(ctx, "presign.bin", 600)
if err != nil {
t.Fatalf("PresignGetURL: %v", err)
}
if !strings.Contains(getURL, "X-Amz-Signature") || !strings.Contains(getURL, "X-Amz-Expires=600") {
t.Fatalf("GET 直链缺少签名参数: %s", getURL)
}
putURL, err := st.PresignPutURL(ctx, "presign.bin", 300)
if err != nil {
t.Fatalf("PresignPutURL: %v", err)
}
if !strings.Contains(putURL, "X-Amz-Signature") {
t.Fatalf("PUT 直链缺少签名参数: %s", putURL)
}
}
// TestS3HealthCheck 健康检查(ListObjectsV2)。
func TestS3HealthCheck(t *testing.T) {
st, _ := newTestS3(t)
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
}
// TestS3RetryOn503 SDK 内置重试器:503 后成功。
func TestS3RetryOn503(t *testing.T) {
st, f := newTestS3(t)
ctx := context.Background()
data := []byte("retry-me")
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "retry.bin"); err != nil {
t.Fatal(err)
}
f.mu.Lock()
f.failNextGet = 1
f.mu.Unlock()
dl, err := st.Open(ctx, "retry.bin", nil)
if err != nil {
t.Fatalf("503 后应重试成功: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("content mismatch")
}
}
// TestS3FactoryRegistry 工厂构造。
func TestS3FactoryRegistry(t *testing.T) {
f := newFakeS3()
srv := httptest.NewServer(f)
defer srv.Close()
prev := engineOptions.S3
engineOptions.S3 = S3Options{
AccessKeyID: "ak", SecretAccessKey: "sk", Bucket: "b",
Endpoint: srv.URL, Region: "us-east-1", AddressingStyle: "path",
}
defer func() { engineOptions.S3 = prev }()
st, err := NewEngine(context.Background(), "s3")
if err != nil {
t.Fatalf("NewEngine(s3): %v", err)
}
if err := st.HealthCheck(context.Background()); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if _, err := st.PresignGetURL(context.Background(), "x.bin", 60); err != nil {
t.Fatalf("Presign: %v", err)
}
_ = time.Now
}
+889
View File
@@ -0,0 +1,889 @@
package storage
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/xml"
"errors"
"fmt"
"io"
"math/rand"
"net"
"net/http"
"net/url"
"os"
"path"
"strconv"
"strings"
"sync"
"time"
)
// WebDAVStorage 基于 net/http 的 WebDAV 引擎(本次重写的重点优化对象)。
//
// 相比参考实现(WebDAVFileStorageaiohttp)的改进:
// - 单例 http.Client + 连接池化 Transport(参考实现每个操作新建 ClientSession,无复用);
// - Basic 与 DigestRFC 2617qop=authMD5/SHA-256)双认证自动协商(参考实现仅 Basic);
// - GET 下载透传 Range 头(参考实现全量 GET,无法断点/分段);
// - 5xx/429/网络错误指数退避重试,可配次数(参考实现无重试);
// - 下载经 io.Pipe 流式转发,全程不落盘;
// - 目录存在性内存缓存,按需逐级 MKCOL,避免每次保存都发 PROPFIND;
// - 非流式操作带可配超时;流式传输由调用方 ctx 管控(可取消)。
type WebDAVStorage struct {
base *url.URL // 服务基址(含可能的路径前缀),以 / 结尾
root string // 远端根目录(webdav_root_path
username string
password string
client *http.Client
transport *http.Transport
auth *authState
maxRetries int // 5xx/网络错误最大重试次数
baseBackoff time.Duration // 退避基数
opTimeout time.Duration // 非流式操作超时
dirMu sync.RWMutex
knownDirs map[string]struct{} // 已确认存在的远端目录(含根前缀)
spacesPool sync.Pool // 256KB 复用缓冲
}
// NewWebDAVStorage 构造 WebDAV 引擎。
func NewWebDAVStorage(opts WebDAVOptions) (*WebDAVStorage, error) {
opts.applyDefaults()
raw := strings.TrimSpace(opts.BaseURL)
if raw == "" {
return nil, fmt.Errorf("storage/webdav: 缺少 webdav_url 配置")
}
if !strings.Contains(raw, "://") {
raw = "http://" + raw
}
u, err := url.Parse(raw)
if err != nil {
return nil, fmt.Errorf("storage/webdav: webdav_url 非法: %w", err)
}
if u.Scheme != "http" && u.Scheme != "https" {
return nil, fmt.Errorf("storage/webdav: webdav_url 仅支持 http/https,收到 %q", u.Scheme)
}
if !strings.HasSuffix(u.Path, "/") {
u.Path += "/"
}
root := strings.Trim(opts.RootPath, "/")
if root == "" {
root = "filebox_storage"
}
root = strings.ReplaceAll(root, "\\", "/")
transport := newPooledTransport(opts.MaxIdleConnsPerHost)
return &WebDAVStorage{
base: u,
root: root,
username: opts.Username,
password: opts.Password,
client: &http.Client{Transport: transport},
transport: transport,
auth: newAuthState(opts.Username, opts.Password),
maxRetries: opts.MaxRetries,
baseBackoff: time.Duration(opts.BaseBackoff) * time.Millisecond,
opTimeout: time.Duration(opts.Timeout) * time.Second,
knownDirs: map[string]struct{}{},
spacesPool: sync.Pool{New: func() any {
b := make([]byte, localChunkSize)
return &b
}},
}, nil
}
func init() {
RegisterEngine("webdav", func(ctx context.Context) (Storage, error) {
return NewWebDAVStorage(engineOptions.WebDAV)
})
}
// newPooledTransport 连接池化 TransportKeep-Alive 连接复用是 WebDAV 优化的核心。
func newPooledTransport(maxIdlePerHost int) *http.Transport {
return &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: 10 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 100,
MaxIdleConnsPerHost: maxIdlePerHost,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: 10 * time.Second,
ExpectContinueTimeout: time.Second,
ResponseHeaderTimeout: 60 * time.Second,
}
}
// requestOpts 单次 WebDAV 请求参数。
type requestOpts struct {
// body 请求体工厂:每次尝试调用一次(重试时重新获取,可重放)。
body func() (io.Reader, int64, error)
// retryBody 请求体是否可重放(seekable);false 时 PUT 类请求失败不重试。
retryBody bool
// headers 附加请求头。
headers map[string]string
// streaming 流式传输(GET/PUT 大 body):不套 opTimeout,由调用方 ctx 管控。
streaming bool
}
// do 执行一次 WebDAV 请求:认证自动协商 + 指数退避重试。
// 返回的响应由调用方负责关闭(drainClose / readErrorBody)。
//
// 重要:非流式操作的可配超时通过 ctx 实现,cancel 不随 do() 返回而调用,
// 而是挂在 davResponse 上、待响应体读完后再触发——否则取消会提前杀掉
// Keep-Alive 连接,破坏连接复用。
func (w *WebDAVStorage) do(ctx context.Context, method, rawURL string, opts requestOpts) (*davResponse, error) {
// 非流式操作套可配超时(流式由调用方 ctx 管控)。
var cancel context.CancelFunc
if !opts.streaming {
if _, hasDeadline := ctx.Deadline(); !hasDeadline {
ctx, cancel = context.WithTimeout(ctx, w.opTimeout)
}
}
fail := func(err error) (*davResponse, error) {
if cancel != nil {
cancel()
}
return nil, err
}
// 幂等方法或可重放 body 才允许整体重试。
idempotent := method == http.MethodGet || method == http.MethodHead ||
method == "PROPFIND" || method == "MKCOL" || method == http.MethodDelete ||
method == http.MethodOptions
retryable := idempotent || opts.retryBody
const maxAuthRetries = 2
budget := w.maxRetries + maxAuthRetries // 认证挑战重试不消耗退避预算
authRetries := 0
for attempt := 0; attempt < budget; attempt++ {
var body io.Reader
var length int64 = -1
if opts.body != nil {
var err error
body, length, err = opts.body()
if err != nil {
return fail(fmt.Errorf("%w: 构造请求体失败: %v", ErrUnavailable, err))
}
}
req, err := http.NewRequestWithContext(ctx, method, rawURL, body)
if err != nil {
return fail(fmt.Errorf("%w: 构造请求失败: %v", ErrInvalidPath, err))
}
if length >= 0 {
req.ContentLength = length
}
for k, v := range opts.headers {
req.Header.Set(k, v)
}
w.auth.apply(req)
resp, err := w.client.Do(req)
if err != nil {
if ctx.Err() != nil { // 调用方取消/超时优先
return fail(ctx.Err())
}
if retryable && attempt+1 < budget {
if sleepErr := w.backoff(ctx, attempt, 0); sleepErr != nil {
return fail(sleepErr)
}
continue
}
return fail(fmt.Errorf("%w: %s %s: %v", ErrUnavailable, method, rawURL, err))
}
// 401 认证挑战:切换 Basic/Digest 后立即重试(不退避、不额外计数)。
if resp.StatusCode == http.StatusUnauthorized && authRetries < maxAuthRetries {
challenge := resp.Header.Get("WWW-Authenticate")
drainClose(&davResponse{Response: resp})
if challenge != "" && w.auth.challenge(challenge) {
authRetries++
continue
}
return fail(fmt.Errorf("%w: WebDAV 认证失败(401%s", ErrUnavailable, rawURL))
}
// 5xx/429/408:幂等或可重放 body 时指数退避重试。
if retryable && isRetryStatus(resp.StatusCode) && attempt+1 < budget {
retryAfter := retryAfterSeconds(resp.Header.Get("Retry-After"))
drainClose(&davResponse{Response: resp})
if sleepErr := w.backoff(ctx, attempt, retryAfter); sleepErr != nil {
return fail(sleepErr)
}
continue
}
return &davResponse{Response: resp, cancel: cancel}, nil
}
return fail(fmt.Errorf("%w: WebDAV 重试耗尽(%s %s", ErrUnavailable, method, rawURL))
}
// davResponse WebDAV 响应 + 关联的超时取消函数。
// 非流式操作读完响应体后必须经 drainClose/readErrorBody 释放(触发 cancel)。
type davResponse struct {
*http.Response
cancel context.CancelFunc
}
// isRetryStatus 判断状态码是否值得重试。
func isRetryStatus(code int) bool {
switch code {
case http.StatusRequestTimeout, http.StatusTooManyRequests,
http.StatusInternalServerError, http.StatusBadGateway,
http.StatusServiceUnavailable, http.StatusGatewayTimeout:
return true
}
return false
}
// retryAfterSeconds 解析 Retry-After(秒);非法或负值返回 0。
func retryAfterSeconds(v string) time.Duration {
if v == "" {
return 0
}
n, err := strconv.Atoi(strings.TrimSpace(v))
if err != nil || n <= 0 {
return 0
}
if n > 5 {
n = 5 // 上限 5s,避免异常服务端拖死请求
}
return time.Duration(n) * time.Second
}
// backoff 指数退避:base * 2^attempt,封顶 2s,带 ±20% 抖动;retryAfter 优先。
func (w *WebDAVStorage) backoff(ctx context.Context, attempt int, retryAfter time.Duration) error {
d := retryAfter
if d <= 0 {
d = w.baseBackoff << attempt
if d > 2*time.Second {
d = 2 * time.Second
}
// ±20% 抖动
jitter := time.Duration(int64(d) / 5)
if jitter > 0 {
d -= time.Duration(rand.Int63n(int64(jitter)))
}
}
if d <= 0 {
d = time.Millisecond
}
timer := time.NewTimer(d)
defer timer.Stop()
select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return nil
}
}
// drainClose 读取少量残余并关闭响应体,保证连接可复用;随后触发超时清理。
func drainClose(resp *davResponse) {
if resp == nil || resp.Body == nil {
return
}
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 8<<10))
_ = resp.Body.Close()
if resp.cancel != nil {
resp.cancel()
}
}
// joinRemote 校验 savePath 并拼接远端完整路径(含根目录前缀)。
func (w *WebDAVStorage) joinRemote(savePath string) (string, error) {
cleaned, ok := SanitizePath(savePath)
if !ok {
return "", fmt.Errorf("%w: %q", ErrInvalidPath, savePath)
}
return path.Join(w.root, cleaned), nil
}
// urlFor 将远端路径转为完整 URLURL.String 自动按段转义)。
func (w *WebDAVStorage) urlFor(remotePath string) string {
u := *w.base
p := strings.TrimSuffix(u.Path, "/")
remotePath = strings.Trim(remotePath, "/")
if remotePath != "" && remotePath != "." {
p += "/" + remotePath
}
u.Path = p
return u.String()
}
// propfindBody PROPFIND 请求体:只取需要的属性。
const propfindBody = `<?xml version="1.0" encoding="utf-8"?>` +
`<D:propfind xmlns:D="DAV:"><D:prop>` +
`<D:resourcetype/><D:getcontentlength/><D:getcontenttype/>` +
`</D:prop></D:propfind>`
// davMultistatus 207 Multi-Status XML 解析结构(标签名与命名空间无关匹配)。
type davMultistatus struct {
Responses []struct {
Href string `xml:"href"`
Propstat []struct {
Status string `xml:"status"`
Prop struct {
ContentLength int64 `xml:"getcontentlength"`
ContentType string `xml:"getcontenttype"`
ResourceType struct {
Collection *struct{} `xml:"collection"`
} `xml:"resourcetype"`
} `xml:"prop"`
} `xml:"propstat"`
} `xml:"response"`
}
// firstProp 取第一个 HTTP 2xx 状态的属性块。
func (m *davMultistatus) firstProp() (length int64, ctype string, isDir bool, ok bool) {
for _, r := range m.Responses {
for _, ps := range r.Propstat {
if !strings.Contains(ps.Status, " 200 ") {
continue
}
return ps.Prop.ContentLength, ps.Prop.ContentType, ps.Prop.ResourceType.Collection != nil, true
}
}
return 0, "", false, false
}
// propfind 执行 PROPFIND 并解析 207 响应;404 时返回 (nil, nil)。
func (w *WebDAVStorage) propfind(ctx context.Context, rawURL string, depth string) (*davMultistatus, error) {
resp, err := w.do(ctx, "PROPFIND", rawURL, requestOpts{
body: func() (io.Reader, int64, error) {
return strings.NewReader(propfindBody), int64(len(propfindBody)), nil
},
headers: map[string]string{"Depth": depth, "Content-Type": "application/xml"},
})
if err != nil {
return nil, err
}
defer drainClose(resp)
switch resp.StatusCode {
case http.StatusMultiStatus, http.StatusOK:
body, err := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
if err != nil {
return nil, fmt.Errorf("%w: PROPFIND 读取失败: %v", ErrUnavailable, err)
}
var ms davMultistatus
if err := xml.Unmarshal(body, &ms); err != nil {
return nil, fmt.Errorf("%w: PROPFIND XML 解析失败: %v", ErrUnavailable, err)
}
return &ms, nil
case http.StatusNotFound:
return nil, nil
default:
return nil, fmt.Errorf("%w: PROPFIND %s → %d", ErrUnavailable, rawURL, resp.StatusCode)
}
}
// remoteExists PROPFIND 探测远端路径存在性。
func (w *WebDAVStorage) remoteExists(ctx context.Context, remotePath string) (bool, error) {
ms, err := w.propfind(ctx, w.urlFor(remotePath), "0")
if err != nil {
return false, err
}
return ms != nil, nil
}
// markDir 记录已确认存在的目录(避免重复 PROPFIND/MKCOL 往返)。
func (w *WebDAVStorage) markDir(remotePath string) {
w.dirMu.Lock()
defer w.dirMu.Unlock()
w.knownDirs[remotePath] = struct{}{}
}
// unmarkDir 目录被删除时移除缓存。
func (w *WebDAVStorage) unmarkDir(remotePath string) {
w.dirMu.Lock()
defer w.dirMu.Unlock()
delete(w.knownDirs, remotePath)
}
// isMarkedDir 查询目录缓存。
func (w *WebDAVStorage) isMarkedDir(remotePath string) bool {
w.dirMu.RLock()
defer w.dirMu.RUnlock()
_, ok := w.knownDirs[remotePath]
return ok
}
// ensureDirs 按需逐级创建远端目录(含根前缀;MKCOL 级联,成功后写缓存)。
func (w *WebDAVStorage) ensureDirs(ctx context.Context, remotePath string) error {
segments := splitRemoteSegments(remotePath)
cur := ""
for _, seg := range segments {
cur = path.Join(cur, seg)
if w.isMarkedDir(cur) {
continue
}
exists, err := w.remoteExists(ctx, cur)
if err == nil && exists {
w.markDir(cur)
continue
}
if err != nil && !errors.Is(err, ErrNotFound) {
return err
}
resp, err := w.do(ctx, "MKCOL", w.urlFor(cur), requestOpts{})
if err != nil {
return err
}
status := resp.StatusCode
drainClose(resp)
// 201 创建成功;405 已存在;其余视为失败(409 通常因父目录缺失,理论上不会出现)。
if status == http.StatusCreated || status == http.StatusOK ||
status == http.StatusNoContent || status == http.StatusMethodNotAllowed {
w.markDir(cur)
continue
}
return fmt.Errorf("%w: MKCOL %s → %d", ErrUnavailable, cur, status)
}
return nil
}
// splitRemoteSegments 拆分远端路径段。
func splitRemoteSegments(p string) []string {
p = strings.Trim(strings.ReplaceAll(p, "\\", "/"), "/")
if p == "" {
return nil
}
return strings.Split(p, "/")
}
// deleteEmptyParents 删除空父目录(含根前缀,但不删根目录本身);尽力而为。
func (w *WebDAVStorage) deleteEmptyParents(ctx context.Context, remotePath string) {
dir := path.Dir(remotePath)
for dir != "" && dir != "." && dir != w.root && strings.HasPrefix(dir+"/", w.root+"/") {
ms, err := w.propfind(ctx, w.urlFor(dir), "1")
if err != nil || ms == nil {
return
}
if len(ms.Responses) > 1 { // 非空(自身 + 子项)
return
}
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(dir), requestOpts{})
if err != nil {
return
}
ok := resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusNoContent
drainClose(resp)
if !ok {
return
}
w.unmarkDir(dir)
dir = path.Dir(dir)
}
}
// putFile PUT 上传:body 工厂每次尝试返回可重放的读取器。
func (w *WebDAVStorage) putFile(ctx context.Context, rawURL string, body func() (io.Reader, int64, error), retryBody bool) (*davResponse, error) {
return w.do(ctx, http.MethodPut, rawURL, requestOpts{
body: body,
retryBody: retryBody,
headers: map[string]string{"Content-Type": "application/octet-stream"},
streaming: true,
})
}
// checkPutStatus 校验 PUT 响应状态。
func checkPutStatus(resp *davResponse, op string) error {
switch resp.StatusCode {
case http.StatusOK, http.StatusCreated, http.StatusNoContent:
drainClose(resp)
return nil
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: %s → %d %s", ErrUnavailable, op, resp.StatusCode, msg)
}
}
// readErrorBody 读取错误响应前 200 字节并释放连接。
func readErrorBody(resp *davResponse) string {
if resp == nil || resp.Body == nil {
return ""
}
b, _ := io.ReadAll(io.LimitReader(resp.Body, 200))
_ = resp.Body.Close()
if resp.cancel != nil {
resp.cancel()
}
return strings.TrimSpace(string(b))
}
// SaveFile 流式保存(PUT):按需建目录,seekable 源可安全重试。
func (w *WebDAVStorage) SaveFile(ctx context.Context, r io.Reader, savePath string) (int64, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return 0, err
}
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
return 0, err
}
// 可重放判定:seekable 源失败后可从头重传(PUT 覆盖语义保证最终一致)。
seeker, seekable := r.(io.Seeker)
var knownLen int64 = -1
if seekable {
if cur, err := seeker.Seek(0, io.SeekCurrent); err == nil {
if end, err := seeker.Seek(0, io.SeekEnd); err == nil {
knownLen = end - cur
_, _ = seeker.Seek(cur, io.SeekStart)
}
}
}
src := &countingReader{r: r}
body := func() (io.Reader, int64, error) {
if seekable {
if _, err := seeker.Seek(0, io.SeekStart); err != nil {
return nil, -1, err
}
src.reset()
}
return src, knownLen, nil
}
resp, err := w.putFile(ctx, w.urlFor(remote), body, seekable)
if err != nil {
return src.count(), err
}
if err := checkPutStatus(resp, fmt.Sprintf("PUT %s", remote)); err != nil {
return src.count(), err
}
return src.count(), nil
}
// DeleteFile DELETE 文件 + 尽力清理空父目录。
func (w *WebDAVStorage) DeleteFile(ctx context.Context, savePath string) error {
remote, err := w.joinRemote(savePath)
if err != nil {
return err
}
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(remote), requestOpts{})
if err != nil {
return err
}
switch resp.StatusCode {
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
drainClose(resp)
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: DELETE %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
}
w.deleteEmptyParents(ctx, remote)
return nil
}
// Open 打开下载流:Range 透传,io.Pipe 流式转发不落盘,ctx 可取消。
func (w *WebDAVStorage) Open(ctx context.Context, savePath string, rng *Range) (*Download, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return nil, err
}
opts := requestOpts{streaming: true}
if rng != nil {
opts.headers = map[string]string{"Range": rangeHeaderValue(rng)}
}
resp, err := w.do(ctx, http.MethodGet, w.urlFor(remote), opts)
if err != nil {
return nil, err
}
switch resp.StatusCode {
case http.StatusOK, http.StatusPartialContent:
// 正常,继续
case http.StatusNotFound:
drainClose(resp)
return nil, ErrNotFound
case http.StatusRequestedRangeNotSatisfiable:
drainClose(resp)
return nil, ErrRangeNotSatisfiable
default:
msg := readErrorBody(resp)
return nil, fmt.Errorf("%w: GET %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
}
total := resp.ContentLength
start, end := int64(0), total-1
if resp.StatusCode == http.StatusPartialContent {
if cr := resp.Header.Get("Content-Range"); cr != "" {
if s0, e0, t0, ok := parseContentRange(cr); ok {
start, end = s0, e0
if t0 >= 0 {
total = t0
}
}
}
}
if total < 0 { // 服务端未给出长度(chunked):按未知大小处理
start, end, total = 0, -1, -1
}
if rng == nil { // 对齐契约:完整文件 Start=0、End=Total-1
start, end = 0, total-1
}
if end < 0 { // 空文件或未知大小:End 未知语义
end = -1
}
// io.Pipe 流式桥接:HTTP 响应体 → 管道 → 调用方,全程不落盘;
// 调用方提前 Close 或 ctx 取消都会终止拷贝并释放连接。
body := resp.Body
pr, pw := io.Pipe()
go func() {
bufp, _ := w.spacesPool.Get().(*[]byte)
_, copyErr := io.CopyBuffer(pw, body, *bufp)
w.spacesPool.Put(bufp)
_ = body.Close()
pw.CloseWithError(copyErr) // copyErr 为 nil 时写入 EOF
}()
context.AfterFunc(ctx, func() {
_ = pw.CloseWithError(ctx.Err())
})
contentType := resp.Header.Get("Content-Type")
return &Download{
ReadCloser: pr,
Start: start,
End: end,
Total: total,
Meta: FileMeta{
Size: total,
ContentType: contentType,
AcceptRanges: true,
},
}, nil
}
// Stat PROPFIND Depth 0 获取元信息。
func (w *WebDAVStorage) Stat(ctx context.Context, savePath string) (*FileMeta, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return nil, err
}
ms, err := w.propfind(ctx, w.urlFor(remote), "0")
if err != nil {
if errors.Is(err, ErrNotFound) {
return nil, ErrNotFound
}
return nil, err
}
if ms == nil {
return nil, ErrNotFound
}
length, ctype, _, ok := ms.firstProp()
if !ok {
return nil, ErrNotFound
}
return &FileMeta{Size: length, ContentType: ctype, AcceptRanges: true}, nil
}
// HeadMeta 读取文件元信息与前 n 字节(WebDAV 实现:PROPFIND + Range GET)。
func (w *WebDAVStorage) HeadMeta(ctx context.Context, savePath string, headBytes int64) (*FileMeta, []byte, error) {
meta, err := w.Stat(ctx, savePath)
if err != nil {
return nil, nil, err
}
if headBytes <= 0 {
return meta, nil, nil
}
dl, err := w.Open(ctx, savePath, &Range{Start: 0, End: headBytes - 1})
if err != nil {
if errors.Is(err, ErrRangeNotSatisfiable) { // 空文件等边界:返回空头
return meta, nil, nil
}
return nil, nil, err
}
defer func() { _ = dl.Close() }()
head := make([]byte, headBytes)
n, _ := io.ReadFull(dl.ReadCloser, head)
return meta, head[:n], nil
}
// SaveChunk 保存分片:落临时文件获得精确长度与可重放 body,PUT 到分片路径。
func (w *WebDAVStorage) SaveChunk(ctx context.Context, uploadID string, chunkIndex int, r io.Reader, savePath string) (int64, error) {
if _, err := w.joinRemote(savePath); err != nil {
return 0, err
}
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, chunkIndex))
if !ok {
return 0, fmt.Errorf("%w: 分片路径非法 %q", ErrInvalidPath, savePath)
}
remote := path.Join(w.root, chunkRel)
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
return 0, err
}
// 分片体积有限(默认 ≤8MB):落临时文件换取精确 Content-Length 与可重试性。
tmp, err := os.CreateTemp("", "fcb-webdav-chunk-*")
if err != nil {
return 0, fmt.Errorf("storage/webdav: 创建临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
size, err := io.CopyBuffer(tmp, r, make([]byte, localChunkSize))
if err != nil {
return 0, fmt.Errorf("storage/webdav: 缓存分片失败: %w", err)
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, fmt.Errorf("storage/webdav: 回卷分片失败: %w", err)
}
body := func() (io.Reader, int64, error) {
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return nil, -1, err
}
return tmp, size, nil
}
resp, err := w.putFile(ctx, w.urlFor(remote), body, true)
if err != nil {
return size, err
}
if err := checkPutStatus(resp, fmt.Sprintf("PUT 分片 %s", remote)); err != nil {
return size, err
}
return size, nil
}
// MergeChunks 合并 WebDAV 分片:
// 逐分片 GET 流式拼入本地临时文件(边拷贝边校验哈希)→ PUT 上传目标 → 清理远端分片与本地临时文件。
// 说明:WebDAV 无服务端聚合能力,合并必须经服务端中转;临时文件仅用于拼接与重试,最终 PUT 可重放。
func (w *WebDAVStorage) MergeChunks(ctx context.Context, uploadID string, total int, verifyHash func(index int) (string, error), savePath string) (int64, string, error) {
if total <= 0 {
return 0, "", fmt.Errorf("storage/webdav: 非法分片总数 %d", total)
}
remote, err := w.joinRemote(savePath)
if err != nil {
return 0, "", err
}
if err := w.ensureDirs(ctx, path.Dir(remote)); err != nil {
return 0, "", err
}
tmp, err := os.CreateTemp("", "fcb-webdav-merge-*")
if err != nil {
return 0, "", fmt.Errorf("storage/webdav: 创建合并临时文件失败: %w", err)
}
tmpName := tmp.Name()
defer func() { _ = tmp.Close(); _ = os.Remove(tmpName) }()
totalHash := sha256.New()
var size int64
for i := 0; i < total; i++ {
if err := ctx.Err(); err != nil {
return 0, "", err
}
chunkRel, ok := SanitizePath(ChunkPartPath(savePath, uploadID, i))
if !ok {
return 0, "", fmt.Errorf("%w: 分片 %d 路径非法", ErrInvalidPath, i)
}
resp, err := w.do(ctx, http.MethodGet, w.urlFor(path.Join(w.root, chunkRel)), requestOpts{streaming: true})
if err != nil {
return 0, "", fmt.Errorf("storage/webdav: 读取分片 %d 失败: %w", i, err)
}
if resp.StatusCode == http.StatusNotFound {
drainClose(resp)
return 0, "", fmt.Errorf("storage/webdav: 分片 %d 不存在", i)
}
if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent {
msg := readErrorBody(resp)
return 0, "", fmt.Errorf("storage/webdav: 读取分片 %d → %d %s", i, resp.StatusCode, msg)
}
chunkHash := sha256.New()
n, err := io.CopyBuffer(io.MultiWriter(tmp, totalHash, chunkHash), resp.Body, make([]byte, localChunkSize))
_ = resp.Body.Close()
if err != nil {
return 0, "", fmt.Errorf("storage/webdav: 拼接分片 %d 失败: %w", i, err)
}
if verifyHash != nil {
expected, err := verifyHash(i)
if err != nil {
return 0, "", err
}
if expected != "" && expected != hex.EncodeToString(chunkHash.Sum(nil)) {
return 0, "", fmt.Errorf("%w: 分片 %d 期望 %s", ErrHashMismatch, i, expected)
}
}
size += n
}
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return 0, "", fmt.Errorf("storage/webdav: 回卷合并文件失败: %w", err)
}
body := func() (io.Reader, int64, error) {
if _, err := tmp.Seek(0, io.SeekStart); err != nil {
return nil, -1, err
}
return tmp, size, nil
}
resp, err := w.putFile(ctx, w.urlFor(remote), body, true)
if err != nil {
return size, "", err
}
if err := checkPutStatus(resp, fmt.Sprintf("PUT 合并 %s", remote)); err != nil {
return size, "", err
}
// 合并成功后清理远端分片目录与本地临时文件(defer 兜底删除本地文件)。
_ = w.CleanChunks(ctx, uploadID, savePath)
return size, hex.EncodeToString(totalHash.Sum(nil)), nil
}
// CleanChunks 递归删除远端分片目录(RFC 4918 DELETE 对 collection 递归)。
func (w *WebDAVStorage) CleanChunks(ctx context.Context, uploadID string, savePath string) error {
dirRel, ok := SanitizePath(chunkDirOf(savePath, uploadID))
if !ok {
return fmt.Errorf("%w: 分片目录非法 %q", ErrInvalidPath, savePath)
}
remote := path.Join(w.root, dirRel)
resp, err := w.do(ctx, http.MethodDelete, w.urlFor(remote), requestOpts{})
if err != nil {
return err
}
switch resp.StatusCode {
case http.StatusOK, http.StatusNoContent, http.StatusNotFound:
drainClose(resp)
w.unmarkDir(remote)
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: 清理分片目录 %s → %d %s", ErrUnavailable, remote, resp.StatusCode, msg)
}
return nil
}
// FileExists PROPFIND 探测存在性;非法路径按不存在处理。
func (w *WebDAVStorage) FileExists(ctx context.Context, savePath string) (bool, error) {
remote, err := w.joinRemote(savePath)
if err != nil {
return false, nil
}
return w.remoteExists(ctx, remote)
}
// PresignGetURL WebDAV 无预签名直链能力。
func (w *WebDAVStorage) PresignGetURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// PresignPutURL WebDAV 无预签名直传能力。
func (w *WebDAVStorage) PresignPutURL(ctx context.Context, savePath string, expires int64) (string, error) {
return "", ErrNotSupported
}
// HealthCheck 健康检查:PROPFIND 根目录;不存在时 MKCOL 创建(启动自愈)。
// 同时完成凭据与连通性验证(do 内 401 协商)。
func (w *WebDAVStorage) HealthCheck(ctx context.Context) error {
exists, err := w.remoteExists(ctx, w.root)
if err == nil && exists {
w.markDir(w.root)
return nil
}
if err != nil && !errors.Is(err, ErrNotFound) {
return fmt.Errorf("%w: WebDAV 健康检查失败: %v", ErrUnavailable, err)
}
resp, err := w.do(ctx, "MKCOL", w.urlFor(w.root), requestOpts{})
if err != nil {
return err
}
switch resp.StatusCode {
case http.StatusCreated, http.StatusOK, http.StatusNoContent, http.StatusMethodNotAllowed:
drainClose(resp)
w.markDir(w.root)
return nil
default:
msg := readErrorBody(resp)
return fmt.Errorf("%w: WebDAV 根目录创建失败 → %d %s", ErrUnavailable, resp.StatusCode, msg)
}
}
// 接口编译期断言。
var _ Storage = (*WebDAVStorage)(nil)
+244
View File
@@ -0,0 +1,244 @@
package storage
import (
"crypto/md5"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"fmt"
"net/http"
"strings"
"sync"
)
// authMode 认证模式(WebDAV 服务端挑战后自动协商)。
type authMode int
const (
authModeUnknown authMode = iota // 未定:先发 Basic 探测
authModeBasic
authModeDigest
)
// authState WebDAV Basic/Digest 认证状态。
//
// 策略:
// - 首个请求预置 Basic;若服务端 401 且挑战为 Digest,则解析挑战参数切换为 Digest;
// - Digest 按 RFC 2617/7616 实现 qop=authMD5 / SHA-256,含 -sess 变体);
// qop 缺失时回退 RFC 2069 旧式响应;
// - nonce 变更时重置 nc 计数;nc/cnonce 在互斥锁内生成保证并发唯一。
type authState struct {
mu sync.Mutex
username string
password string
mode authMode
realm string
nonce string
qop string // 选定的 qop"auth" 或空 = RFC2069
opaque string
algorithm string // MD5 | MD5-sess | SHA-256 | SHA-256-sess
nc uint32
knownBasicOK bool // 已确认 Basic 可用
}
// newAuthState 构造认证状态(默认以 Basic 起步)。
func newAuthState(username, password string) *authState {
return &authState{username: username, password: password}
}
// apply 为请求设置 Authorization 头(每次请求调用,Digest 时消耗一个 nc)。
func (a *authState) apply(req *http.Request) {
a.mu.Lock()
defer a.mu.Unlock()
switch {
case a.mode == authModeDigest && a.nonce != "":
req.Header.Set("Authorization", a.digestHeader(req))
default:
req.SetBasicAuth(a.username, a.password)
}
}
// digestHeader 依据缓存的挑战参数计算 Digest Authorization 头(调用方需持锁)。
func (a *authState) digestHeader(req *http.Request) string {
uri := req.URL.RequestURI()
method := strings.ToUpper(req.Method)
ncStr := fmt.Sprintf("%08x", a.nc+1)
a.nc++
cnonce := randomHex(8)
var ha1 string
switch strings.ToLower(a.algorithm) {
case "md5-sess":
ha1 = hashHex("md5", hashHex("md5", a.username+":"+a.realm+":"+a.password)+":"+a.nonce+":"+cnonce)
case "sha-256-sess":
ha1 = hashHex("sha256", hashHex("sha256", a.username+":"+a.realm+":"+a.password)+":"+a.nonce+":"+cnonce)
case "sha-256":
ha1 = hashHex("sha256", a.username+":"+a.realm+":"+a.password)
default: // md5
ha1 = hashHex("md5", a.username+":"+a.realm+":"+a.password)
}
ha2 := hashHex(algoName(a.algorithm), method+":"+uri)
var response string
var fields []string
esc := escapeDigestValue(a.username)
if a.qop == "" { // RFC 2069
response = hashHex(algoName(a.algorithm), ha1+":"+a.nonce+":"+ha2)
fields = append(fields,
`Digest username="`+esc+`"`,
`realm="`+escapeDigestValue(a.realm)+`"`,
`nonce="`+escapeDigestValue(a.nonce)+`"`,
`uri="`+escapeDigestValue(uri)+`"`,
`response="`+response+`"`)
} else {
response = hashHex(algoName(a.algorithm), ha1+":"+a.nonce+":"+ncStr+":"+cnonce+":"+a.qop+":"+ha2)
fields = append(fields,
`Digest username="`+esc+`"`,
`realm="`+escapeDigestValue(a.realm)+`"`,
`nonce="`+escapeDigestValue(a.nonce)+`"`,
`uri="`+escapeDigestValue(uri)+`"`,
`cnonce="`+cnonce+`"`,
`nc=`+ncStr,
`qop=`+a.qop,
`response="`+response+`"`,
`algorithm=`+a.algorithm)
}
if a.opaque != "" {
fields = append(fields, `opaque="`+escapeDigestValue(a.opaque)+`"`)
}
return strings.Join(fields, ", ")
}
// challenge 处理 401 的 WWW-Authenticate 挑战;返回是否已切换认证方式可重试。
// 返回 false 表示凭据错误或算法不受支持,调用方应直接报错。
func (a *authState) challenge(header string) bool {
a.mu.Lock()
defer a.mu.Unlock()
h := strings.TrimSpace(header)
lower := strings.ToLower(h)
switch {
case strings.HasPrefix(lower, "digest"):
params := parseChallengeParams(strings.TrimPrefix(h[len("Digest"):], " "))
algo := strings.ToUpper(strings.TrimSpace(params["algorithm"]))
if algo == "" {
algo = "MD5"
}
switch algo {
case "MD5", "MD5-SESS", "SHA-256", "SHA-256-SESS":
default:
return false // 不支持的摘要算法
}
if params["nonce"] == "" || params["realm"] == "" {
return false
}
qop := ""
if raw := strings.TrimSpace(params["qop"]); raw != "" {
for _, candidate := range strings.Split(raw, ",") {
if strings.EqualFold(strings.TrimSpace(candidate), "auth") {
qop = "auth"
break
}
}
if qop == "" {
return false // 仅支持 auth-int 等需要 body 哈希的模式
}
}
if a.nonce != params["nonce"] {
a.nc = 0
}
a.realm, a.nonce, a.qop = params["realm"], params["nonce"], qop
a.opaque, a.algorithm = params["opaque"], strings.ToLower(algo)
a.mode = authModeDigest
a.knownBasicOK = false
return true
case strings.HasPrefix(lower, "basic"):
if a.knownBasicOK || a.mode == authModeBasic {
return false // 已用 Basic 仍 401:凭据错误
}
a.mode = authModeBasic
a.knownBasicOK = true
return true
}
return false
}
// parseChallengeParams 解析 "realm=\"x\", nonce=\"y\"" 形式的挑战参数(引号内逗号不切分)。
func parseChallengeParams(s string) map[string]string {
out := map[string]string{}
for _, item := range splitAuthParams(s) {
kv := strings.SplitN(item, "=", 2)
if len(kv) != 2 {
continue
}
k := strings.ToLower(strings.TrimSpace(kv[0]))
v := strings.TrimSpace(kv[1])
if len(v) >= 2 && strings.HasPrefix(v, `"`) && strings.HasSuffix(v, `"`) {
v = v[1 : len(v)-1]
}
out[k] = v
}
return out
}
// splitAuthParams 逗号切分但忽略引号内的逗号。
func splitAuthParams(s string) []string {
var parts []string
var b strings.Builder
inQuote := false
for i := 0; i < len(s); i++ {
c := s[i]
switch {
case c == '"':
inQuote = !inQuote
b.WriteByte(c)
case c == ',' && !inQuote:
if t := strings.TrimSpace(b.String()); t != "" {
parts = append(parts, t)
}
b.Reset()
default:
b.WriteByte(c)
}
}
if t := strings.TrimSpace(b.String()); t != "" {
parts = append(parts, t)
}
return parts
}
// algoName 映射哈希函数名。
func algoName(algorithm string) string {
switch strings.ToLower(algorithm) {
case "sha-256", "sha-256-sess":
return "sha256"
default:
return "md5"
}
}
// hashHex 通用哈希摘要(algo: md5|sha256)。
func hashHex(algo, s string) string {
if algo == "sha256" {
sum := sha256.Sum256([]byte(s))
return hex.EncodeToString(sum[:])
}
sum := md5.Sum([]byte(s))
return hex.EncodeToString(sum[:])
}
// escapeDigestValue 转义引号。
func escapeDigestValue(s string) string {
return strings.ReplaceAll(s, `"`, `\"`)
}
// randomHex 生成 n 字节随机 hex。
func randomHex(n int) string {
b := make([]byte, n)
if _, err := rand.Read(b); err != nil {
// crypto/rand 失败极其罕见;退化为全零仍保持协议可用。
for i := range b {
b[i] = 0
}
}
return hex.EncodeToString(b)
}
+824
View File
@@ -0,0 +1,824 @@
package storage
import (
"bytes"
"context"
"encoding/base64"
"fmt"
"io"
"net"
"net/http"
"net/http/httptest"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
)
// ---- 最小 WebDAV 假服务:PUT/GET/HEAD/PROPFIND/MKCOL/DELETE + Basic/Digest 认证 ----
type davLog struct {
Method string
Path string
Status int
}
type fakeDav struct {
mu sync.Mutex
dirs map[string]bool
files map[string][]byte
// 认证配置:mode = none|basic|digestdigest 配合 algo = MD5|SHA-256。
mode string
username string
password string
realm string
nonce string
opaque string
algo string
failNext map[string]int // method → 剩余 503 次数
logs []davLog
}
func newFakeDav(mode string) *fakeDav {
return &fakeDav{
dirs: map[string]bool{},
files: map[string][]byte{},
mode: mode,
username: "fcb",
password: "fcb-pass",
realm: "test-realm",
nonce: "dcd98b7102dd2f0e8b11d0f600bfb0c0",
opaque: "5ccc069c403ebaf9f0171e9517f40e41",
algo: "MD5",
failNext: map[string]int{},
}
}
// auth 校验请求凭据;失败时写出 401 与对应挑战。
func (f *fakeDav) auth(w http.ResponseWriter, r *http.Request) bool {
if f.mode == "none" {
return true
}
h := r.Header.Get("Authorization")
ok := false
switch f.mode {
case "basic":
ok = h == "Basic "+basicAuth(f.username, f.password)
case "digest":
ok = f.checkDigest(r)
}
if ok {
return true
}
switch f.mode {
case "basic":
w.Header().Set("WWW-Authenticate", `Basic realm="`+f.realm+`"`)
case "digest":
w.Header().Set("WWW-Authenticate", fmt.Sprintf(
`Digest realm="%s", qop="auth", nonce="%s", opaque="%s", algorithm=%s, stale=false`,
f.realm, f.nonce, f.opaque, f.algo))
}
w.WriteHeader(http.StatusUnauthorized)
return false
}
// checkDigest 服务端重算 RFC 2617 摘要响应。
func (f *fakeDav) checkDigest(r *http.Request) bool {
h := r.Header.Get("Authorization")
if !strings.HasPrefix(h, "Digest ") {
return false
}
p := parseChallengeParams(strings.TrimSpace(h[len("Digest "):]))
ha1 := hashHex(algoName(f.algo), f.username+":"+f.realm+":"+f.password)
ha2 := hashHex(algoName(f.algo), strings.ToUpper(r.Method)+":"+r.URL.RequestURI())
got := hashHex(algoName(f.algo), ha1+":"+f.nonce+":"+p["nc"]+":"+p["cnonce"]+":"+p["qop"]+":"+ha2)
return p["username"] == f.username && p["response"] == got
}
func basicAuth(user, pass string) string {
return base64.StdEncoding.EncodeToString([]byte(user + ":" + pass))
}
func (f *fakeDav) record(method, path string, status int) {
f.logs = append(f.logs, davLog{Method: method, Path: path, Status: status})
}
// maybeFail 命中失败注入时返回 true(已写出 503)。
func (f *fakeDav) maybeFail(w http.ResponseWriter, method string) bool {
if f.failNext[method] > 0 {
f.failNext[method]--
w.WriteHeader(http.StatusServiceUnavailable)
return true
}
return false
}
func (f *fakeDav) ServeHTTP(w http.ResponseWriter, r *http.Request) {
f.mu.Lock()
defer f.mu.Unlock()
if !f.auth(w, r) {
f.record(r.Method, r.URL.Path, http.StatusUnauthorized)
return
}
p := strings.Trim(r.URL.Path, "/")
switch r.Method {
case http.MethodPut:
if f.maybeFail(w, "PUT") {
f.record(r.Method, p, 503)
return
}
if parent := parentOf(p); parent != "" && !f.dirs[parent] {
w.WriteHeader(http.StatusConflict) // 强制客户端先建目录
f.record(r.Method, p, 409)
return
}
body, _ := io.ReadAll(r.Body)
f.files[p] = body
f.record(r.Method, p, 201)
w.WriteHeader(http.StatusCreated)
case http.MethodGet:
if f.maybeFail(w, "GET") {
f.record(r.Method, p, 503)
return
}
data, ok := f.files[p]
if !ok {
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
return
}
w.Header().Set("Content-Type", "application/octet-stream")
w.Header().Set("Accept-Ranges", "bytes")
if rng := r.Header.Get("Range"); rng != "" {
start, end := int64(0), int64(len(data))-1
spec := strings.TrimPrefix(rng, "bytes=")
if strings.HasSuffix(spec, "-") { // bytes=N- → 到文件尾
if s, err := strconv.ParseInt(strings.TrimSuffix(spec, "-"), 10, 64); err == nil {
start = s
}
} else if _, err := fmt.Sscanf(spec, "%d-%d", &start, &end); err != nil {
w.WriteHeader(http.StatusBadRequest)
f.record(r.Method, p, 400)
return
}
if start < 0 || start >= int64(len(data)) {
w.WriteHeader(http.StatusRequestedRangeNotSatisfiable)
f.record(r.Method, p, 416)
return
}
if end >= int64(len(data)) {
end = int64(len(data)) - 1
}
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(data)))
w.WriteHeader(http.StatusPartialContent)
_, _ = w.Write(data[start : end+1])
f.record(r.Method, p, 206)
return
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write(data)
f.record(r.Method, p, 200)
case http.MethodHead:
data, ok := f.files[p]
if !ok {
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
return
}
w.Header().Set("Content-Length", strconv.Itoa(len(data)))
w.WriteHeader(http.StatusOK)
f.record(r.Method, p, 200)
case "PROPFIND":
depth := r.Header.Get("Depth")
self, isDirSelf := f.stat(p)
if !isDirSelf {
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
return
}
var b strings.Builder
b.WriteString(`<?xml version="1.0" encoding="utf-8"?>` +
`<D:multistatus xmlns:D="DAV:">`)
f.writeResponse(&b, p, self)
if depth == "1" && self.isDir {
for _, name := range f.children(p) {
child := name
cs, cd := f.stat(child)
f.writeResponse(&b, child, davStat{isDir: cd, size: cs.size})
}
}
b.WriteString(`</D:multistatus>`)
w.Header().Set("Content-Type", "application/xml; charset=utf-8")
w.WriteHeader(http.StatusMultiStatus)
_, _ = w.Write([]byte(b.String()))
f.record(r.Method, p, 207)
case "MKCOL":
if f.dirs[p] || f.files[p] != nil {
w.WriteHeader(http.StatusMethodNotAllowed) // 已存在
f.record(r.Method, p, 405)
return
}
if parent := parentOf(p); parent != "" && !f.dirs[parent] {
w.WriteHeader(http.StatusConflict)
f.record(r.Method, p, 409)
return
}
f.dirs[p] = true
w.WriteHeader(http.StatusCreated)
f.record(r.Method, p, 201)
case http.MethodDelete:
if _, ok := f.files[p]; ok {
delete(f.files, p)
w.WriteHeader(http.StatusNoContent)
f.record(r.Method, p, 204)
return
}
if f.dirs[p] {
// 递归删除目录
prefix := p + "/"
for name := range f.files {
if strings.HasPrefix(name, prefix) {
delete(f.files, name)
}
}
for name := range f.dirs {
if name == p || strings.HasPrefix(name+"/", prefix) {
delete(f.dirs, name)
}
}
w.WriteHeader(http.StatusNoContent)
f.record(r.Method, p, 204)
return
}
w.WriteHeader(http.StatusNotFound)
f.record(r.Method, p, 404)
default:
w.WriteHeader(http.StatusMethodNotAllowed)
f.record(r.Method, p, 405)
}
}
type davStat struct {
isDir bool
size int
}
func (f *fakeDav) stat(p string) (davStat, bool) {
if data, ok := f.files[p]; ok {
return davStat{size: len(data)}, true
}
if f.dirs[p] {
return davStat{isDir: true}, true
}
return davStat{}, false
}
func (f *fakeDav) children(p string) []string {
var out []string
prefix := p + "/"
for name := range f.files {
if strings.HasPrefix(name, prefix) && !strings.Contains(strings.TrimPrefix(name, prefix), "/") {
out = append(out, name)
}
}
for name := range f.dirs {
if strings.HasPrefix(name, prefix) && !strings.Contains(strings.TrimPrefix(name, prefix), "/") {
out = append(out, name)
}
}
return out
}
func (f *fakeDav) writeResponse(b *strings.Builder, href string, st davStat) {
b.WriteString(`<D:response><D:href>/` + href + `</D:href><D:propstat><D:prop><D:resourcetype>`)
if st.isDir {
b.WriteString(`<D:collection/>`)
}
b.WriteString(`</D:resourcetype><D:getcontentlength>` + strconv.Itoa(st.size) +
`</D:getcontentlength><D:getcontenttype>application/octet-stream</D:getcontenttype>` +
`</D:prop><D:status>HTTP/1.1 200 OK</D:status></D:propstat></D:response>`)
}
func parentOf(p string) string {
if i := strings.LastIndex(p, "/"); i > 0 {
return p[:i]
}
return ""
}
// newTestDav 构造 WebDAV 引擎 + 假服务。
func newTestDav(t *testing.T, mode string, tweak func(o *WebDAVOptions)) (*WebDAVStorage, *fakeDav, *int32) {
t.Helper()
f := newFakeDav(mode)
var conns int32
srv := httptest.NewUnstartedServer(f)
srv.Config.ConnState = func(c net.Conn, cs http.ConnState) {
if cs == http.StateNew {
atomic.AddInt32(&conns, 1)
}
}
srv.Start()
t.Cleanup(srv.Close)
opts := WebDAVOptions{
BaseURL: srv.URL,
Username: f.username,
Password: f.password,
RootPath: "fcb_root",
MaxRetries: 3,
}
if tweak != nil {
tweak(&opts)
}
st, err := NewWebDAVStorage(opts)
if err != nil {
t.Fatalf("NewWebDAVStorage: %v", err)
}
return st, f, &conns
}
// TestWebDAVBasicCRUD Basic 认证下的完整 CRUD 与 Range。
func TestWebDAVBasicCRUD(t *testing.T) {
st, f, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
// 健康检查:根目录 404 → MKCOL 自建
if err := st.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if !f.dirs["fcb_root"] {
t.Fatalf("根目录应被自动创建")
}
data := []byte("WebDAV 引擎数据 0123456789 ABCDEF")
n, err := st.SaveFile(ctx, bytes.NewReader(data), "2025/08/w.bin")
if err != nil {
t.Fatalf("SaveFile: %v", err)
}
if n != int64(len(data)) {
t.Fatalf("n = %d", n)
}
if string(f.files["fcb_root/2025/08/w.bin"]) != string(data) {
t.Fatalf("PUT 内容不匹配")
}
// 按需建目录:两级目录都应已创建
if !f.dirs["fcb_root/2025"] || !f.dirs["fcb_root/2025/08"] {
t.Fatalf("目录未按需创建: %v %v", f.dirs["fcb_root/2025"], f.dirs["fcb_root/2025/08"])
}
meta, err := st.Stat(ctx, "2025/08/w.bin")
if err != nil {
t.Fatalf("Stat: %v", err)
}
if meta.Size != int64(len(data)) || !meta.AcceptRanges {
t.Fatalf("Stat = %+v", meta)
}
if ok, _ := st.FileExists(ctx, "2025/08/w.bin"); !ok {
t.Fatalf("FileExists 应为 true")
}
// 完整下载(对齐 go-api 约定)
dl, err := st.Open(ctx, "2025/08/w.bin", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, err := io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read: %v", err)
}
if !bytes.Equal(got, data) {
t.Fatalf("full mismatch")
}
if dl.Start != 0 || dl.End != int64(len(data))-1 || dl.Total != int64(len(data)) {
t.Fatalf("full offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// Range 下载
dl, err = st.Open(ctx, "2025/08/w.bin", &Range{Start: 2, End: 7})
if err != nil {
t.Fatalf("Open range: %v", err)
}
got, err = io.ReadAll(dl)
_ = dl.Close()
if err != nil {
t.Fatalf("read range: %v", err)
}
if !bytes.Equal(got, data[2:8]) {
t.Fatalf("range mismatch")
}
if dl.Start != 2 || dl.End != 7 || dl.Total != int64(len(data)) {
t.Fatalf("range offsets = %d,%d,%d", dl.Start, dl.End, dl.Total)
}
// 416(起点越界)/ 404
if _, err := st.Open(ctx, "2025/08/w.bin", &Range{Start: int64(len(data)) + 9, End: -1}); err == nil ||
!strings.Contains(err.Error(), ErrRangeNotSatisfiable.Error()) {
t.Fatalf("want ErrRangeNotSatisfiable, got %v", err)
}
if _, err := st.Open(ctx, "no/such.bin", nil); err == nil ||
!strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("want ErrNotFound, got %v", err)
}
// 删除 + 空父目录清理
if err := st.DeleteFile(ctx, "2025/08/w.bin"); err != nil {
t.Fatalf("DeleteFile: %v", err)
}
if ok, _ := st.FileExists(ctx, "2025/08/w.bin"); ok {
t.Fatalf("删除后仍存在")
}
if _, err := st.Stat(ctx, "2025/08/w.bin"); err == nil ||
!strings.Contains(err.Error(), ErrNotFound.Error()) {
t.Fatalf("Stat 应 ErrNotFound, got %v", err)
}
}
// TestWebDAVDigestAuth DigestMD5)认证协商。
func TestWebDAVDigestAuth(t *testing.T) {
st, f, _ := newTestDav(t, "digest", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck(digest): %v", err)
}
// HealthCheck 流程应观察到 401 挑战(客户端先 Basic 探测 → 401 → Digest 重试)
saw401 := false
for _, l := range f.logs {
if l.Status == 401 {
saw401 = true
}
}
if !saw401 {
t.Fatalf("未观察到 401 挑战: %+v", f.logs)
}
// 认证后的 PROPFINDStat 已有目录)应得到 207
if _, err := st.Stat(ctx, ""); err == nil {
// Stat("") 非法路径属预期;这里换用 FileExists 对已有根目录探测
_ = err
}
data := []byte("digest 内容")
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "d.bin"); err != nil {
t.Fatalf("SaveFile(digest): %v", err)
}
dl, err := st.Open(ctx, "d.bin", nil)
if err != nil {
t.Fatalf("Open(digest): %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("digest 下载内容不匹配")
}
// 全链路完成:确认存在成功的 2xx/207 请求
saw2xx := false
for _, l := range f.logs {
if l.Status == 207 || l.Status == 201 || l.Status == 200 {
saw2xx = true
}
}
if !saw2xx {
t.Fatalf("认证后应有成功请求: %+v", f.logs)
}
}
// TestWebDAVDigestSHA256 DigestSHA-256)算法。
func TestWebDAVDigestSHA256(t *testing.T) {
st, _, _ := newTestDav(t, "digest", nil)
st.auth.mu.Lock()
st.auth.algorithm = "sha-256"
st.auth.mu.Unlock()
// 服务端也切换到 SHA-256 重算摘要
st2, f, _ := newTestDav(t, "digest", nil)
f.algo = "SHA-256"
// 先让客户端完成一次 MD5 协商拿到挑战参数,再切 SHA-256 会 401 失败——
// 因此这里直接对 SHA-256 服务端做完整链路(client 首次探测 Basic→401→Digest)。
_ = st
ctx := context.Background()
if err := st2.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck(SHA-256): %v", err)
}
}
// TestWebDAVDigestWrongPassword 凭据错误 → 明确报错而非重试风暴。
func TestWebDAVDigestWrongPassword(t *testing.T) {
f := newFakeDav("digest")
srv := httptest.NewServer(f)
defer srv.Close()
st, err := NewWebDAVStorage(WebDAVOptions{
BaseURL: srv.URL, Username: f.username, Password: "WRONG",
RootPath: "r", MaxRetries: 1, BaseBackoff: 1,
})
if err != nil {
t.Fatal(err)
}
if err := st.HealthCheck(context.Background()); err == nil ||
!strings.Contains(err.Error(), "401") {
t.Fatalf("错误凭据应报 401 相关错误, got %v", err)
}
}
// TestWebDAVRetryGet 5xx 指数退避重试(GET 幂等)。
func TestWebDAVRetryGet(t *testing.T) {
st, f, _ := newTestDav(t, "basic", func(o *WebDAVOptions) { o.BaseBackoff = 5 })
ctx := context.Background()
data := []byte("retry target")
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "r.bin"); err != nil {
t.Fatal(err)
}
f.mu.Lock()
f.failNext["GET"] = 2
f.mu.Unlock()
dl, err := st.Open(ctx, "r.bin", nil)
if err != nil {
t.Fatalf("503×2 后应重试成功: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if !bytes.Equal(got, data) {
t.Fatalf("content mismatch")
}
// 验证确实发了 3 次 GET
gets := 0
for _, l := range f.logs {
if l.Method == "GET" && strings.HasSuffix(l.Path, "r.bin") {
gets++
}
}
if gets != 3 {
t.Fatalf("GET 次数 = %d, want 3", gets)
}
}
// TestWebDAVRetryPut 可重放 bodyseekable)PUT 失败重试;不可重放不重试。
func TestWebDAVRetryPut(t *testing.T) {
st, f, _ := newTestDav(t, "basic", func(o *WebDAVOptions) { o.BaseBackoff = 5 })
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
// seekable:重试成功
f.mu.Lock()
f.failNext["PUT"] = 1
f.mu.Unlock()
data := []byte("put with retry")
if _, err := st.SaveFile(ctx, bytes.NewReader(data), "pr.bin"); err != nil {
t.Fatalf("PUT 重试应成功: %v", err)
}
puts := 0
for _, l := range f.logs {
if l.Method == "PUT" && strings.HasSuffix(l.Path, "pr.bin") {
puts++
}
}
if puts != 2 {
t.Fatalf("PUT 次数 = %d, want 2", puts)
}
// 非 seekableio.Pipe):不重试,直接失败
f.mu.Lock()
f.failNext["PUT"] = 1
f.mu.Unlock()
pr, pw := io.Pipe()
go func() {
_, _ = pw.Write([]byte("non-seekable"))
_ = pw.Close()
}()
if _, err := st.SaveFile(ctx, pr, "ns.bin"); err == nil {
t.Fatalf("非重放 PUT 注入 503 应失败")
}
}
// TestWebDAVConnectionReuse 连接复用:多次请求不应各建一条 TCP 连接。
func TestWebDAVConnectionReuse(t *testing.T) {
st, _, conns := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
for i := 0; i < 12; i++ {
if _, err := st.SaveFile(ctx, bytes.NewReader([]byte("x")), fmt.Sprintf("reuse/%d.bin", i)); err != nil {
t.Fatal(err)
}
if _, err := st.Stat(ctx, fmt.Sprintf("reuse/%d.bin", i)); err != nil {
t.Fatal(err)
}
}
// 25 次请求(12 PUT + 12 PROPFIND + 1 HealthCheck 的 PROPFIND/MKCOL)只允许极少量新连接
if got := atomic.LoadInt32(conns); got > 4 {
t.Fatalf("新建 TCP 连接数 = %d,连接复用失效(应 ≤4)", got)
}
}
// TestWebDAVPipeStreaming io.Pipe 流式转发:完整读取 + 提前关闭。
func TestWebDAVPipeStreaming(t *testing.T) {
st, _, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
big := bytes.Repeat([]byte("0123456789abcdef"), 4096) // 64KB
if _, err := st.SaveFile(ctx, bytes.NewReader(big), "big.bin"); err != nil {
t.Fatal(err)
}
dl, err := st.Open(ctx, "big.bin", nil)
if err != nil {
t.Fatal(err)
}
got, err := io.ReadAll(dl)
if err != nil {
t.Fatalf("read pipe: %v", err)
}
_ = dl.Close()
if !bytes.Equal(got, big) {
t.Fatalf("pipe content mismatch")
}
// 提前关闭:后续读取返回错误且不挂死
dl2, err := st.Open(ctx, "big.bin", nil)
if err != nil {
t.Fatal(err)
}
buf := make([]byte, 10)
if _, err := io.ReadFull(dl2, buf); err != nil {
t.Fatalf("read head: %v", err)
}
if err := dl2.Close(); err != nil {
t.Fatalf("early close: %v", err)
}
// ctx 取消同样会终止流
cctx, cancel := context.WithCancel(context.Background())
dl3, err := st.Open(cctx, "big.bin", nil)
if err != nil {
t.Fatal(err)
}
cancel()
time.Sleep(20 * time.Millisecond)
_, err = dl3.Read(buf)
if err == nil {
_ = dl3.Close()
t.Fatalf("ctx 取消后读取应报错")
}
_ = dl3.Close()
}
// TestWebDAVChunkMerge 分片保存/合并/清理。
func TestWebDAVChunkMerge(t *testing.T) {
st, f, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
savePath := "2025/09/merged.bin"
uploadID := "uid-webdav"
chunks := [][]byte{[]byte("AAA"), []byte("BB"), []byte("CCCC")}
hashes := make([]string, len(chunks))
for i, c := range chunks {
n, err := st.SaveChunk(ctx, uploadID, i, bytes.NewReader(c), savePath)
if err != nil {
t.Fatalf("SaveChunk %d: %v", i, err)
}
if n != int64(len(c)) {
t.Fatalf("chunk %d size = %d", i, n)
}
hashes[i] = sha256Hex(c)
}
size, fileHash, err := st.MergeChunks(ctx, uploadID, len(chunks), func(i int) (string, error) {
return hashes[i], nil
}, savePath)
if err != nil {
t.Fatalf("MergeChunks: %v", err)
}
if size != 9 || fileHash != sha256Hex(bytes.Join(chunks, nil)) {
t.Fatalf("merge result = %d %s", size, fileHash)
}
if string(f.files["fcb_root/"+savePath]) != "AAABBCCCC" {
t.Fatalf("合并内容错误: %q", f.files["fcb_root/"+savePath])
}
// 分片目录已清理
for k := range f.files {
if strings.Contains(k, "chunks/"+uploadID) {
t.Fatalf("分片残留: %s", k)
}
}
if f.dirs["fcb_root/2025/09/chunks/"+uploadID] {
t.Fatalf("分片目录残留")
}
}
// TestWebDAVCleanChunks 清理与哈希失败路径。
func TestWebDAVCleanChunks(t *testing.T) {
st, f, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatal(err)
}
for i := 0; i < 2; i++ {
if _, err := st.SaveChunk(ctx, "uidc", i, strings.NewReader("z"), "c.bin"); err != nil {
t.Fatal(err)
}
}
if err := st.CleanChunks(ctx, "uidc", "c.bin"); err != nil {
t.Fatalf("CleanChunks: %v", err)
}
if len(f.files) != 0 {
t.Fatalf("分片未清理: %v", f.files)
}
// 幂等
if err := st.CleanChunks(ctx, "uidc", "c.bin"); err != nil {
t.Fatalf("CleanChunks idempotent: %v", err)
}
// 哈希不匹配
if _, err := st.SaveChunk(ctx, "uidm", 0, strings.NewReader("real"), "m.bin"); err != nil {
t.Fatal(err)
}
if _, _, err := st.MergeChunks(ctx, "uidm", 1, func(i int) (string, error) {
return sha256Hex([]byte("wrong")), nil
}, "m.bin"); err == nil || !strings.Contains(err.Error(), ErrHashMismatch.Error()) {
t.Fatalf("want ErrHashMismatch, got %v", err)
}
}
// TestWebDAVTimeout 非流式操作超时:PROPFIND 响应慢于 Timeout1s)→ context deadline exceeded。
func TestWebDAVTimeout(t *testing.T) {
f := newFakeDav("basic")
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method == "PROPFIND" {
time.Sleep(1500 * time.Millisecond) // > Timeout 1s
}
f.ServeHTTP(w, r)
}))
defer srv.Close()
st, err := NewWebDAVStorage(WebDAVOptions{
BaseURL: srv.URL, Username: f.username, Password: f.password,
RootPath: "r", MaxRetries: 0, Timeout: 1, BaseBackoff: 5,
})
if err != nil {
t.Fatal(err)
}
start := time.Now()
err = st.HealthCheck(context.Background())
if err == nil {
t.Fatalf("超时应报错")
}
if !strings.Contains(err.Error(), "context deadline exceeded") {
t.Fatalf("应为超时错误, got %v", err)
}
// 单次尝试 1s 超时 + 一次重试 ≈ 2s;若超时未生效会拖满 2×1.5s
if elapsed := time.Since(start); elapsed > 3500*time.Millisecond {
t.Fatalf("超时未生效(耗时 %v", elapsed)
}
}
// TestWebDAVPresignNotSupported 预签名 → ErrNotSupported。
func TestWebDAVPresignNotSupported(t *testing.T) {
st, _, _ := newTestDav(t, "basic", nil)
ctx := context.Background()
if _, err := st.PresignGetURL(ctx, "x.bin", 60); err == nil ||
!strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("want ErrNotSupported, got %v", err)
}
if _, err := st.PresignPutURL(ctx, "x.bin", 60); err == nil ||
!strings.Contains(err.Error(), ErrNotSupported.Error()) {
t.Fatalf("want ErrNotSupported, got %v", err)
}
}
// TestWebDAVFactoryRegistry 工厂构造 + Digest 全链路。
func TestWebDAVFactoryRegistry(t *testing.T) {
f := newFakeDav("digest")
srv := httptest.NewServer(f)
defer srv.Close()
prev := engineOptions.WebDAV
engineOptions.WebDAV = WebDAVOptions{
BaseURL: srv.URL, Username: f.username, Password: f.password,
RootPath: "factory_root", MaxRetries: 3, BaseBackoff: 5,
}
defer func() { engineOptions.WebDAV = prev }()
st, err := NewEngine(context.Background(), "webdav")
if err != nil {
t.Fatalf("NewEngine(webdav): %v", err)
}
ctx := context.Background()
if err := st.HealthCheck(ctx); err != nil {
t.Fatalf("HealthCheck: %v", err)
}
if _, err := st.SaveFile(ctx, strings.NewReader("factory"), "f.txt"); err != nil {
t.Fatalf("SaveFile: %v", err)
}
dl, err := st.Open(ctx, "f.txt", nil)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, _ := io.ReadAll(dl)
_ = dl.Close()
if string(got) != "factory" {
t.Fatalf("content = %q", got)
}
}