FileCodeBox Go 重写版 v2.5.6(安全审计修复版)
Go 1.27.1 (Gin+GORM) + Vue 3 文件快传服务: - 安全审计全部修复(docs/security-audit-2026-09-05.md): bcrypt 密码哈希与自动升级、presign 直传服务端大小/内容校验、 全局请求体上限、依赖升级(govulncheck 0 命中)、janitor 后台清理、 管理端审计动作落库、/admin CORS 收紧、通知内容白名单净化、 会话默认 7 天、限流缓存故障降级、robots.txt 端点等 - 前端:取件链接复制修复(不再重复拼接提取码)、markdown 净化器加固 - Redis 支持库号(FCB_REDIS_DB / redis://…/db URL) - 文档:docs/api/* 与 openapi.yaml 同步最新行为(robots.txt、 提码 5 位起、chunk 32MiB 上限、admin 审计动作等) 验证:gofmt/go vet/go test 全绿;二进制端到端冒烟通过
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -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 单分片大小上限 32MB(M3:限制 io.ReadAll 内存占用)。
|
||||
const maxChunkSizeBytes = 32 * 1024 * 1024
|
||||
|
||||
// ============ POST /chunk/upload/init 初始化分片会话 ============
|
||||
|
||||
// requireChunkEnabled L4:enableChunk 开关后端强制(此前仅前端隐藏入口,
|
||||
// 开关关闭后 /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_size,0=回落 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/upload,upload_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": "上传已取消"})
|
||||
}
|
||||
@@ -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_type(secret/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 生成下载令牌(L2:HMAC-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_count(0=不限制,超限 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())
|
||||
}
|
||||
// 原子条件插入:已用 + 生效预留 + 本次 <= 上限。
|
||||
// Info5:Postgres 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 兼容归一化:旧前端 bundle(fetch 字符串 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]
|
||||
}
|
||||
@@ -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 应报错")
|
||||
}
|
||||
// day:7 天内合法
|
||||
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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
// policy.go — v2 上传策略统一读取与校验(需求 ④⑩)。
|
||||
//
|
||||
// 管理端在后台设置页修改策略(settings KV,t1 schema)后,上传链路
|
||||
// (share/file、chunk、presign)每次请求实时读取当前策略并动态校验:
|
||||
// - 大小上限:max_file_size(0=回落 uploadSize,语义见 config.MaxFileSize);
|
||||
// - 类型白名单:allowed_file_types("*" 不限制),由 validateFileMagic 统一执行;
|
||||
// - 保存策略:expire_style 白名单 + max_save_seconds 时间上限 + max_save_count
|
||||
// 次数上限,统一在 resolveExpire(helpers.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
|
||||
}
|
||||
@@ -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 构造带真实依赖的 Deps:sqlite 文件库(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:包装为 Manager(build 直接返回 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'}
|
||||
|
||||
// ============ ① 公开 config:v2 展示与策略字段下发 ============
|
||||
|
||||
// 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/update:v2 键全链路 + 类型范围校验 ============
|
||||
|
||||
// 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
|
||||
}
|
||||
@@ -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_size,0=回落 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))
|
||||
}
|
||||
@@ -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)
|
||||
// Info3:robotsText 配置键此前无路由承接,补上(内容可在管理端自定义)
|
||||
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_secret,settings.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 token(403)。
|
||||
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
|
||||
}
|
||||
@@ -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 L4:enableChunk=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 M3:chunk_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 L3:4 位自定义码拒绝、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("错误密码不应通过")
|
||||
}
|
||||
}
|
||||
@@ -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("&", "&", "<", "<", ">", ">", `"`, "'")
|
||||
return r.Replace(s)
|
||||
}
|
||||
@@ -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(配合全局 BodyLimit;441KB 为 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_size,0=回落 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
|
||||
}
|
||||
@@ -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 表单调用 handler(v3.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 body(fetch 字符串 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())
|
||||
}
|
||||
}
|
||||
@@ -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.html(SPA 回退用)。
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user