347 lines
6.6 KiB
Go
347 lines
6.6 KiB
Go
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
|
||
}
|