Files
2026-08-23 00:48:10 +08:00

424 lines
11 KiB
Go

package service
import (
"crypto/md5"
"encoding/hex"
"fmt"
"io"
"mime"
"mime/multipart"
"os"
"path/filepath"
"strings"
"time"
"github.com/gin-gonic/gin"
"filestoragesystem/internal/config"
"filestoragesystem/internal/model"
"filestoragesystem/internal/repository"
"filestoragesystem/internal/utils"
"filestoragesystem/pkg/apperr"
)
// FileService 文件服务
type FileService struct {
cfg *config.Config
fileRepo *repository.FileRepo
projectRepo *repository.ProjectRepo
userRepo *repository.UserRepo
tempLinkRepo *repository.TempLinkRepo
trafficRepo *repository.TrafficRepo
settingRepo *repository.SettingRepo
opLogRepo *repository.OpLogRepo
webhookSvc *WebhookService
}
// settingInt 读取数字型设置
func (s *FileService) settingInt(key string, def int64) int64 {
if v, err := s.settingRepo.GetValue(key); err == nil {
var n int64
if _, e := fmt.Sscanf(v, "%d", &n); e == nil && n > 0 {
return n
}
}
return def
}
// settingBool 读取布尔型设置
func (s *FileService) settingBool(key string, def bool) bool {
if v, err := s.settingRepo.GetValue(key); err == nil {
return v == "true" || v == "1"
}
return def
}
// Upload 上传文件:校验配额/类型/大小,MD5去重,落盘,记录统计
func (s *FileService) Upload(c *gin.Context, userID uint, fh *multipart.FileHeader) (*model.File, error) {
projectID, _ := strconvUint(c.PostForm("project_id"))
if projectID == 0 {
return nil, fmt.Errorf("缺少项目ID")
}
project, err := s.projectRepo.FindByID(projectID)
if err != nil || project.UserID != userID {
return nil, apperr.ErrForbidden
}
if project.Status != 1 {
return nil, fmt.Errorf("项目已被禁用")
}
src, err := fh.Open()
if err != nil {
return nil, err
}
defer src.Close()
size := fh.Size
// 大小限制
maxSize := s.settingInt("max_file_size", 104857600)
if size > maxSize {
return nil, apperr.ErrFileTooLarge
}
// 类型限制
allowed := "*"
if v, err := s.settingRepo.GetValue("allowed_file_types"); err == nil {
allowed = v
}
if !s.checkTypeAllowed(fh.Filename, allowed) {
return nil, apperr.ErrFileTypeDenied
}
// 配额
if project.StorageUsed+size > project.StorageLimit {
return nil, apperr.ErrQuotaExceeded
}
user, err := s.userRepo.FindByID(userID)
if err != nil {
return nil, apperr.ErrNotFound
}
if user.StorageUsed+size > user.StorageLimit {
return nil, apperr.ErrQuotaExceeded
}
// 虚拟路径 + 文件名
virtualPath := utils.SanitizePath(c.PostForm("path"))
baseName := filepath.Base(utils.SanitizePath(fh.Filename))
if baseName == "" || baseName == "." || baseName == "/" {
baseName = "unnamed"
}
filename := baseName
if virtualPath != "" {
filename = virtualPath + "/" + baseName
}
visibility := int8(1)
if c.PostForm("visibility") == "2" || c.PostForm("visibility") == "public" {
visibility = 2
}
// 写入临时文件并计算MD5
tmpFile, err := os.CreateTemp("", "fss-upload-*")
if err != nil {
return nil, err
}
tmpPath := tmpFile.Name()
defer os.Remove(tmpPath)
hasher := md5.New()
if _, err := io.Copy(io.MultiWriter(tmpFile, hasher), src); err != nil {
tmpFile.Close()
return nil, err
}
tmpFile.Close()
md5sum := hex.EncodeToString(hasher.Sum(nil))
// 去重:相同MD5复用已有物理文件
storedRel := ""
if exist, err := s.fileRepo.FindByMD5AndPath(md5sum); err == nil && exist.StoredPath != "" {
if _, serr := os.Stat(filepath.Join(s.cfg.Storage.Root, exist.StoredPath)); serr == nil {
storedRel = exist.StoredPath
}
}
// 无可复用文件则保存到 blob/<md5前2位>/<md5>
if storedRel == "" {
blobRel := filepath.ToSlash(filepath.Join("blobs", md5sum[:2], md5sum))
blobAbs := filepath.Join(s.cfg.Storage.Root, blobRel)
if err := os.MkdirAll(filepath.Dir(blobAbs), 0o755); err != nil {
return nil, err
}
if err := os.Rename(tmpPath, blobAbs); err != nil {
// 跨盘时回退为复制
if cerr := copyFile(tmpPath, blobAbs); cerr != nil {
return nil, cerr
}
}
storedRel = blobRel
}
mimeType := fh.Header.Get("Content-Type")
if mimeType == "" || mimeType == "application/octet-stream" {
if t := mime.TypeByExtension(strings.ToLower(filepath.Ext(baseName))); t != "" {
mimeType = t
}
}
f := &model.File{
ProjectID: projectID,
UserID: userID,
Filename: filename,
StoredPath: storedRel,
Size: size,
MimeType: mimeType,
MD5: md5sum,
Visibility: visibility,
Status: 1,
}
if err := s.fileRepo.Create(f); err != nil {
return nil, err
}
// 更新配额、流量、webhook
_ = s.projectRepo.UpdateStorageUsed(projectID, size)
_ = s.userRepo.UpdateStorageUsed(userID, size)
_ = s.trafficRepo.Create(&model.TrafficLog{
UserID: userID, ProjectID: projectID, FileID: f.ID,
Type: 1, Size: size, IP: utils.ClientIP(c), CreatedAt: time.Now(),
})
s.webhookSvc.Dispatch(userID, projectID, "file.upload", map[string]interface{}{
"file_id": f.ID, "filename": f.Filename, "size": f.Size, "md5": f.MD5, "project_id": projectID,
})
utils.Info("storage", "用户#%d 上传文件 %s (%d字节)", userID, f.Filename, f.Size)
return f, nil
}
func (s *FileService) checkTypeAllowed(filename, allowed string) bool {
allowed = strings.TrimSpace(allowed)
if allowed == "" || allowed == "*" {
return true
}
ext := strings.ToLower(strings.TrimPrefix(filepath.Ext(filename), "."))
if ext == "" {
return true // 无扩展名不限制
}
for _, t := range strings.Split(allowed, ",") {
t = strings.ToLower(strings.TrimSpace(strings.TrimPrefix(t, ".")))
if t == "*" || t == ext {
return true
}
}
return false
}
// GetInfo 获取文件信息(所有者或管理员)
func (s *FileService) GetInfo(userID uint, id uint, isAdmin bool) (*model.File, error) {
f, err := s.fileRepo.FindByID(id)
if err != nil {
return nil, apperr.ErrNotFound
}
if !isAdmin && f.UserID != userID {
return nil, apperr.ErrForbidden
}
return f, nil
}
// List 文件列表
func (s *FileService) List(userID uint, isAdmin bool, filter repository.FileListFilter) ([]model.File, int64, error) {
if !isAdmin {
filter.UserID = userID
}
if filter.Page <= 0 {
filter.Page = 1
}
if filter.PageSize <= 0 {
filter.PageSize = 20
}
return s.fileRepo.List(filter)
}
// Open 打开文件物理句柄并校验访问权,记录流量
func (s *FileService) Open(c *gin.Context, userID uint, id uint, isAdmin bool) (*model.File, *os.File, error) {
f, err := s.GetInfo(userID, id, isAdmin)
if err != nil {
return nil, nil, err
}
if f.Status != 1 {
return nil, nil, apperr.ErrFileDeleted
}
abs := filepath.Join(s.cfg.Storage.Root, f.StoredPath)
fp, err := os.Open(abs)
if err != nil {
utils.Error("storage", "打开文件失败 path=%s err=%v", abs, err)
return nil, nil, fmt.Errorf("文件读取失败")
}
return f, fp, nil
}
// RecordDownload 记录下载行为
func (s *FileService) RecordDownload(c *gin.Context, f *model.File) {
_ = s.fileRepo.IncrDownloadCount(f.ID)
_ = s.trafficRepo.Create(&model.TrafficLog{
UserID: f.UserID, ProjectID: f.ProjectID, FileID: f.ID,
Type: 2, Size: f.Size, IP: utils.ClientIP(c), CreatedAt: time.Now(),
})
}
// Delete 删除文件;permanent=true时硬删除
func (s *FileService) Delete(userID uint, id uint, isAdmin bool, permanent bool) error {
f, err := s.GetInfo(userID, id, isAdmin)
if err != nil {
return err
}
if !permanent {
if f.Status == 0 {
return apperr.ErrFileDeleted
}
if err := s.fileRepo.SoftDelete(id); err != nil {
return err
}
// 回收站不占配额
_ = s.projectRepo.UpdateStorageUsed(f.ProjectID, -f.Size)
_ = s.userRepo.UpdateStorageUsed(f.UserID, -f.Size)
return nil
}
// 硬删除
if err := s.fileRepo.Delete(id); err != nil {
return err
}
_ = s.tempLinkRepo.DeleteByFile(id)
// 若处于回收站,配额已扣减过
if f.Status == 1 {
_ = s.projectRepo.UpdateStorageUsed(f.ProjectID, -f.Size)
_ = s.userRepo.UpdateStorageUsed(f.UserID, -f.Size)
}
// 物理文件无其他引用时删除
if n, err := s.fileRepo.CountByStoredPath(f.StoredPath); err == nil && n == 0 {
_ = os.Remove(filepath.Join(s.cfg.Storage.Root, f.StoredPath))
}
s.webhookSvc.Dispatch(f.UserID, f.ProjectID, "file.delete", map[string]interface{}{
"file_id": f.ID, "filename": f.Filename, "permanent": permanent,
})
return nil
}
// Restore 从回收站恢复文件
func (s *FileService) Restore(userID uint, id uint, isAdmin bool) error {
f, err := s.GetInfo(userID, id, isAdmin)
if err != nil {
return err
}
if f.Status != 0 {
return fmt.Errorf("文件不在回收站中")
}
// 恢复前检查配额
project, err := s.projectRepo.FindByID(f.ProjectID)
if err == nil {
if project.StorageUsed+f.Size > project.StorageLimit {
return apperr.ErrQuotaExceeded
}
}
user, err := s.userRepo.FindByID(f.UserID)
if err == nil && user.StorageUsed+f.Size > user.StorageLimit {
return apperr.ErrQuotaExceeded
}
if err := s.fileRepo.Restore(id); err != nil {
return err
}
_ = s.projectRepo.UpdateStorageUsed(f.ProjectID, f.Size)
_ = s.userRepo.UpdateStorageUsed(f.UserID, f.Size)
return nil
}
// CreateTempLink 生成临时访问链接
func (s *FileService) CreateTempLink(userID uint, id uint, isAdmin bool, password string, maxCount int, expiresInSeconds int64) (*model.TempLink, error) {
f, err := s.GetInfo(userID, id, isAdmin)
if err != nil {
return nil, err
}
if f.Status != 1 {
return nil, apperr.ErrFileDeleted
}
if expiresInSeconds <= 0 {
expiresInSeconds = s.settingInt("temp_link_max_age", 3600)
}
if expiresInSeconds > 7*24*3600 {
expiresInSeconds = 7 * 24 * 3600
}
link := &model.TempLink{
FileID: f.ID,
UserID: userID,
Token: utils.RandomKey(48),
MaxCount: maxCount,
ExpiresAt: time.Now().Add(time.Duration(expiresInSeconds) * time.Second),
}
if password != "" {
link.Password = password
}
if err := s.tempLinkRepo.Create(link); err != nil {
return nil, err
}
link.HasPwd = link.Password != ""
return link, nil
}
// ResolveTempLink 通过token解析临时链接(校验有效期/次数/密码),返回文件与物理句柄
func (s *FileService) ResolveTempLink(token, password string) (*model.File, *os.File, error) {
link, err := s.tempLinkRepo.FindByToken(token)
if err != nil {
return nil, nil, apperr.ErrLinkExpired
}
if time.Now().After(link.ExpiresAt) {
return nil, nil, apperr.ErrLinkExpired
}
if link.MaxCount > 0 && link.UsedCount >= link.MaxCount {
return nil, nil, apperr.ErrLinkExpired
}
if link.Password != "" && link.Password != password {
return nil, nil, apperr.ErrLinkPassword
}
f, err := s.fileRepo.FindByID(link.FileID)
if err != nil || f.Status != 1 {
return nil, nil, apperr.ErrNotFound
}
fp, err := os.Open(filepath.Join(s.cfg.Storage.Root, f.StoredPath))
if err != nil {
return nil, nil, fmt.Errorf("文件读取失败")
}
_ = s.tempLinkRepo.IncrUsed(link.ID)
return f, fp, nil
}
// ---- 辅助函数 ----
func copyFile(src, dst string) error {
in, err := os.Open(src)
if err != nil {
return err
}
defer in.Close()
out, err := os.Create(dst)
if err != nil {
return err
}
defer out.Close()
if _, err := io.Copy(out, in); err != nil {
return err
}
return out.Sync()
}
func strconvUint(s string) (uint, error) {
var n uint
if s == "" {
return 0, fmt.Errorf("empty")
}
for _, ch := range s {
if ch < '0' || ch > '9' {
return 0, fmt.Errorf("invalid")
}
n = n*10 + uint(ch-'0')
}
return n, nil
}