Files
agent/backend/internal/logic/files/service.go

175 lines
5.0 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package files
import (
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"time"
"gorm.io/gorm"
"senlinai-agent/backend/internal/models"
)
type Service struct {
root string
db *gorm.DB
createTemp func(string, string) (stagedFile, error)
rename func(string, string) error
remove func(string) error
}
type stagedFile interface {
io.Writer
Close() error
Name() string
}
type StoredFile struct {
OriginalName string
RelativePath string
AbsolutePath string
}
// NewService 集中持有存储根目录;任何服务端本地路径都只能由该服务构造。
func NewService(root string, databases ...*gorm.DB) *Service {
var database *gorm.DB
if len(databases) > 0 {
database = databases[0]
}
return &Service{
root: root, db: database,
createTemp: func(directory, pattern string) (stagedFile, error) { return os.CreateTemp(directory, pattern) },
rename: os.Rename,
remove: os.Remove,
}
}
// Save 清理客户端文件名,并只返回供持久化的相对路径;绝对路径不得进入 API DTO。
func (s *Service) Save(projectID uint, originalName string, content io.Reader) (StoredFile, error) {
cleanName := filepath.Base(strings.ReplaceAll(strings.TrimSpace(originalName), "\\", "/"))
if cleanName == "." || cleanName == "" {
cleanName = "upload.bin"
}
relative := filepath.ToSlash(filepath.Join("projects", fmt.Sprint(projectID), fmt.Sprintf("%d-%s", time.Now().UnixNano(), cleanName)))
absolute, err := s.absolutePath(relative)
if err != nil {
return StoredFile{}, err
}
if err := os.MkdirAll(filepath.Dir(absolute), 0o755); err != nil {
return StoredFile{}, err
}
// 临时文件与最终文件位于同一目录Close 成功后再原子替换,避免暴露半写入内容。
file, err := s.createTemp(filepath.Dir(absolute), ".upload-*")
if err != nil {
return StoredFile{}, err
}
if _, err := io.Copy(file, content); err != nil {
_ = file.Close()
s.cleanupFailedSave(file.Name(), absolute)
return StoredFile{}, err
}
if err := file.Close(); err != nil {
s.cleanupFailedSave(file.Name(), absolute)
return StoredFile{}, err
}
if err := s.rename(file.Name(), absolute); err != nil {
s.cleanupFailedSave(file.Name(), absolute)
return StoredFile{}, err
}
return StoredFile{OriginalName: cleanName, RelativePath: relative, AbsolutePath: absolute}, nil
}
// Remove 只依据 Service 生成的相对路径定位文件,并尽量移除直至存储根目录的空父目录。
func (s *Service) Remove(stored StoredFile) error {
absolute, err := s.absolutePath(stored.RelativePath)
if err != nil {
return err
}
if err := s.remove(absolute); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
root, err := filepath.Abs(s.root)
if err != nil {
return err
}
for directory := filepath.Dir(absolute); directory != root; directory = filepath.Dir(directory) {
entries, err := os.ReadDir(directory)
if errors.Is(err, os.ErrNotExist) {
continue
}
if err != nil {
return err
}
if len(entries) > 0 {
return nil
}
if err := s.remove(directory); err != nil && !errors.Is(err, os.ErrNotExist) {
return err
}
}
return nil
}
func (s *Service) absolutePath(relativePath string) (string, error) {
cleaned := filepath.Clean(filepath.FromSlash(strings.TrimSpace(relativePath)))
if cleaned == "." || filepath.IsAbs(cleaned) || cleaned == ".." || strings.HasPrefix(cleaned, ".."+string(filepath.Separator)) {
return "", ErrSourcePathRequired
}
root, err := filepath.Abs(s.root)
if err != nil {
return "", err
}
absolute, err := filepath.Abs(filepath.Join(root, cleaned))
if err != nil {
return "", err
}
relativeToRoot, err := filepath.Rel(root, absolute)
if err != nil || relativeToRoot == ".." || strings.HasPrefix(relativeToRoot, ".."+string(filepath.Separator)) {
return "", ErrSourcePathRequired
}
return absolute, nil
}
func (s *Service) cleanupFailedSave(tempPath, finalPath string) {
_ = s.remove(tempPath)
_ = s.remove(finalPath)
}
func (s *Service) database() *gorm.DB {
if s.db != nil {
return s.db
}
return models.DBService
}
var (
ErrSourceTitleRequired = errors.New("source title is required")
ErrSourcePathRequired = errors.New("source file path is required")
)
// CreateSource 只持久化 Save 产生的相对路径,不接受 handler 自行拼接本地路径。
func (s *Service) CreateSource(ownerID uint, project *models.SenlinAgentProject, title string, stored StoredFile) (*models.SenlinAgentSource, error) {
title = strings.TrimSpace(title)
if title == "" {
title = stored.OriginalName
}
if title == "" {
return nil, ErrSourceTitleRequired
}
relativePath := filepath.ToSlash(strings.TrimSpace(stored.RelativePath))
if relativePath == "" || filepath.IsAbs(relativePath) {
return nil, ErrSourcePathRequired
}
source := &models.SenlinAgentSource{
ProjectID: project.ID, ProjectIdentity: project.Identity, CreatedBy: ownerID,
Kind: "file", Title: title, FilePath: relativePath,
}
if err := s.database().Create(source).Error; err != nil {
return nil, err
}
return source, nil
}