package api import ( "errors" "fmt" "io" "net/http" "strconv" "strings" "time" "github.com/gin-gonic/gin" "gorm.io/gorm" "fileshare/internal/audit" "fileshare/internal/middleware" "fileshare/internal/model" "fileshare/internal/response" "fileshare/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) // v3.2:下载带宽限速(storage.ReadCloser → 限速 reader → c.Writer) dlReader := middleware.WrapReadCloser(dl, d.Cfg.DownloadRate()) n, _ := io.Copy(c.Writer, dlReader) 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 }