CI 测试 / go vet + go test (push) Successful in 49s
- 全项目版本号统一:v3.x 迭代号(26.9/26.9/26.9/26.9 及裸 v2/v3)→ 26.9, 覆盖 Go 注释 / 文档 / openapi.yaml / README×4 / 前端源码(80+ 处) - v31_test.go 更名 custom_code_test.go;TestV2AccessorDefaults → TestKVAccessorDefaults - docs/api/00-overview.md 更新日志合并为单条 26.9 条目(修复错位拼接) - .goreleaser.yaml 头部注释与实际一致(Pro 2.18.1 / GITEA_TOKEN / semver tag 要求) - CI:release-image.yml → ci.yml,仅保留 vet+test 门禁; 镜像发布移交 GoReleaser Pro(原 build-push 的 tag 校验与 26.9 版本方案冲突,历史 9 次失败) - 前端重建:server/web/dist 与 web-embed 同步(docs 文案嵌入更新)
799 lines
28 KiB
Go
799 lines
28 KiB
Go
// 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"
|
||
|
||
"fileshare/internal/audit"
|
||
"fileshare/internal/config"
|
||
"fileshare/internal/middleware"
|
||
"fileshare/internal/model"
|
||
"fileshare/internal/response"
|
||
"fileshare/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)
|
||
}
|
||
|
||
// ============ 自定义提取码(26.9,防撞库)============
|
||
|
||
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 26.9:分享记录创建失败时,若是自定义码唯一索引冲突(并发兜底,
|
||
// 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
|
||
}
|
||
|
||
// ============ 站点对外域名(26.9)============
|
||
|
||
// SiteDomain 站点对外域名规范化(26.9):空串合法(分享链接用当前访问地址)。
|
||
// 接受 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 天上限;
|
||
// 26.9 需求 ④: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 26.9:按归属引擎取存储实例(空戳/未知名回落当前引擎,兼容历史数据)。
|
||
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 自适应;
|
||
// 26.9 需求④: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
|
||
}
|
||
// 26.9 兼容归一化:旧前端 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]
|
||
}
|