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// 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 }