package admin import ( "context" "image" "io" "mime/multipart" "net/http" "os" "path/filepath" "server/internal/config" "server/internal/db" "server/internal/db/sqlc" "server/internal/model/common" "server/internal/model/request" "server/internal/pkg/errs" "server/internal/pkg/httputil" "server/internal/pkg/safego" "strings" "sync" _ "image/jpeg" _ "image/png" gonanoid "github.com/matoous/go-nanoid/v2" ) type FileService struct { store *db.Store config *config.Config } func NewFileService(store *db.Store, config *config.Config) *FileService { return &FileService{ store: store, config: config, } } // MakeSavedDir 创建目录并返回 func MakeSavedDir(uploadDir, folder string) (string, error) { rootDir, err := os.Getwd() if err != nil { return "", err } dir := filepath.Join(rootDir, uploadDir, folder) if err = os.MkdirAll(dir, 0755); err != nil { return "", err } return dir, nil } func MakeFilePath(rootDir string, uploadDir string, filePath string) string { return filepath.Join(rootDir, uploadDir, filePath) } // isWebp 判断是否为 WebP 魔数(RIFF+WEBP),DetectContentType 不识别需手动补 func isWebp(b []byte) bool { return len(b) >= 12 && string(b[0:4]) == "RIFF" && string(b[8:12]) == "WEBP" } // detectMime 检测文件类型 因为从header里获取的可能是伪造的 func detectMime(file multipart.File) (string, error) { buf := make([]byte, 512) n, err := file.Read(buf) if err != nil { return "", err } mimeType := http.DetectContentType(buf[:n]) // 补识别 DetectContentType 漏掉的格式 if mimeType == "application/octet-stream" && isWebp(buf[:n]) { mimeType = "image/webp" } // 回到文件开头 if _, err := file.Seek(0, io.SeekStart); err != nil { return "", err } return mimeType, nil } func (s *FileService) List(ctx context.Context, p *common.Pagination) (*common.PageResult[sqlc.ListFilesRow], error) { params := sqlc.ListFilesParams{ Limit: p.PageSize, Offset: (p.Page - 1) * p.PageSize, } total, err := s.store.CountFiles(ctx) if err != nil { return nil, err } list, err := s.store.ListFiles(ctx, params) if err != nil { return nil, err } return &common.PageResult[sqlc.ListFilesRow]{ List: list, Total: total, }, nil } // sanitizeOriginalName 过滤文件名中的 < > 和控制字符,防 HTML 注入 func sanitizeOriginalName(name string) string { var b strings.Builder b.Grow(len(name)) for _, r := range name { switch { case r == '<' || r == '>': continue case r < 0x20 || r == 0x7f: continue // 控制字符 default: b.WriteRune(r) } } // 限长 255 个字符 runes := []rune(b.String()) if len(runes) > 255 { runes = runes[:255] } return string(runes) } func (s *FileService) Upload(ctx context.Context, folder string, file *multipart.FileHeader) (*sqlc.CreateFileRow, error) { // 生成文件名 fileID, err := gonanoid.New() if err != nil { return nil, err } savedDir, err := MakeSavedDir(s.config.File.UploadDir, folder) if err != nil { return nil, err } fileExt := filepath.Ext(file.Filename) filename := fileID + fileExt filePath := filepath.Join(folder, filename) // 用于保存文件 savedPath := filepath.Join(savedDir, filename) // 打开上传的文件 src, err := file.Open() if err != nil { return nil, err } defer src.Close() // 检测类型 mimeType, err := detectMime(src) if err != nil { return nil, err } // 内容与扩展名必须匹配,防改后缀绕过 if !request.MatchUploadMime(file.Filename, mimeType) { return nil, errs.ErrFileTypeNotAllowed } var imageMeta *sqlc.CreateFileImageMetadataParams // 图片需能解码,否则拒绝 if strings.HasPrefix(mimeType, "image/") { cfg, format, err := image.DecodeConfig(src) if err != nil { // webp 标准库无法解码,但 MIME 已验证,跳过元数据 if mimeType != "image/webp" { return nil, errs.ErrFileTypeNotAllowed } } else { imageMeta = &sqlc.CreateFileImageMetadataParams{ Width: int32(cfg.Width), Height: int32(cfg.Height), Format: format, } } // 重置读取位置 _, err = src.Seek(0, io.SeekStart) if err != nil { return nil, err } } // 创建目标文件 dst, err := os.Create(savedPath) if err != nil { return nil, err } defer dst.Close() // 复制文件内容 if _, err = dst.ReadFrom(src); err != nil { return nil, err } // 净化文件名,全被过滤则用存储名兜底 originalName := sanitizeOriginalName(file.Filename) if originalName == "" { originalName = filename } params := sqlc.CreateFileParams{ FileName: filename, FilePath: filePath, FileUrl: httputil.BuildFileUrl(&filePath), OriginalName: originalName, FolderName: folder, MimeType: mimeType, FileSize: file.Size, } result, err := db.WithTxResult[*sqlc.CreateFileRow](ctx, s.store, func(q *sqlc.Queries) (*sqlc.CreateFileRow, error) { result, err := q.CreateFile(ctx, params) if err != nil { return nil, err } if imageMeta != nil { meta := *imageMeta meta.FileID = result.ID if err = q.CreateFileImageMetadata(ctx, meta); err != nil { return nil, err } } return &result, nil }) // 事务失败 清理已保存的文件 if err != nil { _ = os.Remove(savedPath) return nil, err } return result, nil } func (s *FileService) SyncMetadata(ctx context.Context) error { files, err := s.store.ListImageFiles(ctx) if err != nil { return err } if len(files) == 0 { return nil } rootDir, err := os.Getwd() if err != nil { return err } var ( metadataList []sqlc.CopyFileImageMetadataParams wg sync.WaitGroup mu sync.Mutex ) // 限制并发数量 sem := make(chan struct{}, 10) for _, f := range files { wg.Add(1) safego.Go(func() { defer wg.Done() sem <- struct{}{} defer func() { <-sem }() absolutePath := MakeFilePath( rootDir, s.config.File.UploadDir, f.FilePath, ) src, err := os.Open(absolutePath) if err != nil { return } defer src.Close() cfg, format, err := image.DecodeConfig(src) if err != nil { return } item := sqlc.CopyFileImageMetadataParams{ FileID: f.ID, Width: int32(cfg.Width), Height: int32(cfg.Height), Format: format, } mu.Lock() metadataList = append(metadataList, item) mu.Unlock() }) } wg.Wait() if len(metadataList) == 0 { return nil } err = s.store.WithTx(ctx, func(q *sqlc.Queries) error { if err = q.TruncateFileImageMetadata(ctx); err != nil { return err } if _, err = q.CopyFileImageMetadata(ctx, metadataList); err != nil { return err } return nil }) if err != nil { return err } return nil }