Refactor file upload handling in neo package
- Simplified the file header conversion process by introducing GetHeader function to streamline header management during uploads. - Updated handleUpload method to utilize the new GetHeader function, improving code clarity and reducing error handling complexity. - Enhanced test cases to reflect changes in upload logic and ensure accurate content validation. - Improved chunked upload handling by implementing a more robust mechanism for tracking upload progress and managing chunk data.
This commit is contained in:
parent
5b99f3bf44
commit
2222bbfcd7
5 changed files with 172 additions and 67 deletions
|
|
@ -248,15 +248,8 @@ func (neo *DSL) handleUpload(c *gin.Context) {
|
|||
os.Remove(file.Filename)
|
||||
}()
|
||||
|
||||
// Convert the header to a FileHeader
|
||||
header, err := attachment.ToFileHeader(file.Header)
|
||||
if err != nil {
|
||||
c.JSON(400, gin.H{"message": err.Error(), "code": 400})
|
||||
c.Done()
|
||||
return
|
||||
}
|
||||
|
||||
// Upload the file
|
||||
header := attachment.GetHeader(c.Request.Header, file.Header)
|
||||
res, err := manager.Upload(c.Request.Context(), header, reader, option)
|
||||
if err != nil {
|
||||
c.JSON(400, gin.H{"message": err.Error(), "code": 500})
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
package local
|
||||
|
||||
import (
|
||||
"compress/gzip"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
|
|
@ -142,10 +143,21 @@ func (storage *Storage) MergeChunks(ctx context.Context, fileID string, totalChu
|
|||
// Reader read file from local storage
|
||||
func (storage *Storage) Reader(ctx context.Context, fileID string) (io.ReadCloser, error) {
|
||||
fullpath := filepath.Join(storage.Path, fileID)
|
||||
|
||||
reader, err := os.Open(fullpath)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// If the file is a gzip file, decompress it
|
||||
if strings.HasSuffix(fileID, ".gz") {
|
||||
reader, err := gzip.NewReader(reader)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return reader, nil
|
||||
}
|
||||
|
||||
return reader, nil
|
||||
}
|
||||
|
||||
|
|
@ -159,7 +171,7 @@ func (storage *Storage) Download(ctx context.Context, fileID string) (io.ReadClo
|
|||
|
||||
// Try to detect content type from file extension
|
||||
contentType := "application/octet-stream"
|
||||
ext := filepath.Ext(path)
|
||||
ext := filepath.Ext(strings.TrimSuffix(fileID, ".gz"))
|
||||
switch strings.ToLower(ext) {
|
||||
case ".txt":
|
||||
contentType = "text/plain"
|
||||
|
|
@ -179,6 +191,28 @@ func (storage *Storage) Download(ctx context.Context, fileID string) (io.ReadClo
|
|||
contentType = "image/gif"
|
||||
case ".pdf":
|
||||
contentType = "application/pdf"
|
||||
case ".mp4":
|
||||
contentType = "video/mp4"
|
||||
case ".mp3":
|
||||
contentType = "audio/mpeg"
|
||||
case ".wav":
|
||||
contentType = "audio/wav"
|
||||
case ".ogg":
|
||||
contentType = "audio/ogg"
|
||||
case ".webm":
|
||||
contentType = "video/webm"
|
||||
case ".webp":
|
||||
contentType = "image/webp"
|
||||
case ".zip":
|
||||
}
|
||||
|
||||
// If the file is a gzip file, decompress it
|
||||
if strings.HasSuffix(fileID, ".gz") {
|
||||
reader, err := gzip.NewReader(reader)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
return reader, contentType, nil
|
||||
}
|
||||
|
||||
return reader, contentType, nil
|
||||
|
|
|
|||
|
|
@ -9,11 +9,13 @@ import (
|
|||
"io"
|
||||
"mime"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"net/textproto"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/yaoapp/yao/config"
|
||||
|
|
@ -23,6 +25,45 @@ import (
|
|||
|
||||
// Managers the managers
|
||||
var Managers = map[string]*Manager{}
|
||||
var uploadChunks = sync.Map{}
|
||||
|
||||
// UploadChunk is the chunk data
|
||||
type UploadChunk struct {
|
||||
Last int
|
||||
Total int64
|
||||
Chunksize int64
|
||||
TotalChunks int64
|
||||
}
|
||||
|
||||
// GetHeader gets the header from the file header and request header
|
||||
func GetHeader(requestHeader http.Header, fileHeader textproto.MIMEHeader) *FileHeader {
|
||||
|
||||
// Convert the header to a FileHeader
|
||||
header := &FileHeader{FileHeader: &multipart.FileHeader{Header: make(map[string][]string)}}
|
||||
|
||||
for key, values := range fileHeader {
|
||||
for _, value := range values {
|
||||
header.Header.Set(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
// Set Content-Sync, Content-Uid, Content-Range
|
||||
if requestHeader.Get("Content-Sync") != "" {
|
||||
header.Header.Set("Content-Sync", requestHeader.Get("Content-Sync"))
|
||||
}
|
||||
|
||||
// Set Content-Uid
|
||||
if requestHeader.Get("Content-Uid") != "" {
|
||||
header.Header.Set("Content-Uid", requestHeader.Get("Content-Uid"))
|
||||
}
|
||||
|
||||
// Set Content-Range
|
||||
if requestHeader.Get("Content-Range") != "" {
|
||||
header.Header.Set("Content-Range", requestHeader.Get("Content-Range"))
|
||||
}
|
||||
|
||||
return header
|
||||
}
|
||||
|
||||
// Register registers a global attachment manager
|
||||
func Register(name string, driver string, option ManagerOption) (*Manager, error) {
|
||||
|
|
@ -38,25 +79,6 @@ func Register(name string, driver string, option ManagerOption) (*Manager, error
|
|||
return manager, nil
|
||||
}
|
||||
|
||||
// ToFileHeader converts a multipart.FileHeader or textproto.MIMEHeader to a FileHeader
|
||||
func ToFileHeader(header interface{}) (*FileHeader, error) {
|
||||
|
||||
switch header := header.(type) {
|
||||
case *multipart.FileHeader:
|
||||
return &FileHeader{
|
||||
FileHeader: header,
|
||||
}, nil
|
||||
case textproto.MIMEHeader:
|
||||
return &FileHeader{
|
||||
FileHeader: &multipart.FileHeader{
|
||||
Header: header,
|
||||
},
|
||||
}, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("invalid header type: %T", header)
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterDefault registers a default attachment manager
|
||||
func RegisterDefault(name string) (*Manager, error) {
|
||||
|
||||
|
|
@ -186,7 +208,7 @@ func New(option ManagerOption) (*Manager, error) {
|
|||
return manager, nil
|
||||
}
|
||||
|
||||
// Upload uploads a file
|
||||
// Upload uploads a file, Content-Sync must be true for chunked upload
|
||||
func (manager Manager) Upload(ctx context.Context, fileheader *FileHeader, reader io.Reader, option UploadOption) (*File, error) {
|
||||
|
||||
file, err := manager.makeFile(fileheader, option)
|
||||
|
|
@ -196,36 +218,38 @@ func (manager Manager) Upload(ctx context.Context, fileheader *FileHeader, reade
|
|||
|
||||
// Handle chunked upload
|
||||
if fileheader.IsChunk() {
|
||||
start, _, total, err := fileheader.GetChunkInfo()
|
||||
start, end, total, err := fileheader.GetChunkInfo()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid chunk info: %w", err)
|
||||
}
|
||||
|
||||
// Calculate chunk index based on start position and a standard chunk size
|
||||
// We need to determine the standard chunk size from the first few chunks
|
||||
standardChunkSize := int64(1024) // Default chunk size
|
||||
if start > 0 {
|
||||
// For non-first chunks, we can infer the standard chunk size
|
||||
// by looking at the start position
|
||||
if start%1024 == 0 {
|
||||
standardChunkSize = 1024
|
||||
} else if start%2048 == 0 {
|
||||
standardChunkSize = 2048
|
||||
} else if start%4096 == 0 {
|
||||
standardChunkSize = 4096
|
||||
} else {
|
||||
// Try to infer from the start position
|
||||
for size := int64(512); size <= 8192; size *= 2 {
|
||||
if start%size == 0 {
|
||||
standardChunkSize = size
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
// Store the chunk info
|
||||
chunkIndex := 0
|
||||
if start == 0 {
|
||||
chunksize := end - start + 1
|
||||
totalChunks := (total + chunksize - 1) / chunksize
|
||||
uploadChunks.LoadOrStore(file.ID, &UploadChunk{
|
||||
Last: chunkIndex,
|
||||
Total: total,
|
||||
Chunksize: chunksize,
|
||||
TotalChunks: totalChunks,
|
||||
})
|
||||
}
|
||||
|
||||
chunkIndex := int(start / standardChunkSize)
|
||||
totalChunks := int((total + standardChunkSize - 1) / standardChunkSize) // Ceiling division
|
||||
// Update the chunk index
|
||||
v, ok := uploadChunks.Load(file.ID)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("chunk data not found")
|
||||
}
|
||||
|
||||
chunkdata := v.(*UploadChunk)
|
||||
|
||||
// Update the chunk index
|
||||
if start != 0 {
|
||||
chunkIndex = chunkdata.Last + 1
|
||||
chunkdata.Last = chunkIndex
|
||||
uploadChunks.Store(file.ID, chunkdata)
|
||||
}
|
||||
|
||||
// Apply gzip compression if requested
|
||||
if option.Gzip {
|
||||
|
|
@ -234,6 +258,7 @@ func (manager Manager) Upload(ctx context.Context, fileheader *FileHeader, reade
|
|||
return nil, fmt.Errorf("failed to gzip chunk: %w", err)
|
||||
}
|
||||
reader = bytes.NewReader(compressed)
|
||||
|
||||
}
|
||||
|
||||
// Upload chunk
|
||||
|
|
@ -244,7 +269,7 @@ func (manager Manager) Upload(ctx context.Context, fileheader *FileHeader, reade
|
|||
|
||||
// If this is the last chunk, merge all chunks
|
||||
if fileheader.Complete() {
|
||||
err = manager.storage.MergeChunks(ctx, file.ID, totalChunks)
|
||||
err = manager.storage.MergeChunks(ctx, file.ID, int(chunkdata.TotalChunks))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
|
@ -256,6 +281,9 @@ func (manager Manager) Upload(ctx context.Context, fileheader *FileHeader, reade
|
|||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// Remove the chunk data
|
||||
uploadChunks.Delete(file.ID)
|
||||
}
|
||||
|
||||
return file, nil
|
||||
|
|
@ -514,7 +542,11 @@ func (manager Manager) generateFileID(file *FileHeader, extension string, option
|
|||
path = filepath.Join(path, option.AssistantID)
|
||||
}
|
||||
|
||||
return filepath.Join(path, hash[:2], hash[2:4], hash) + extension, nil
|
||||
id := filepath.Join(path, hash[:2], hash[2:4], hash) + extension
|
||||
if option.Gzip {
|
||||
id = id + ".gz"
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// getSize converts the size to bytes
|
||||
|
|
|
|||
|
|
@ -118,14 +118,8 @@ func TestManagerUpload(t *testing.T) {
|
|||
t.Fatalf("Failed to read gzipped file: %v", err)
|
||||
}
|
||||
|
||||
// Since we're storing compressed data, we need to decompress it
|
||||
decompressed, err := Gunzip(downloadedContent)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to decompress file: %v", err)
|
||||
}
|
||||
|
||||
if string(decompressed) != content {
|
||||
t.Errorf("Expected content '%s', got '%s'", content, string(decompressed))
|
||||
if string(downloadedContent) != content {
|
||||
t.Errorf("Expected content '%s', got '%s'", content, string(downloadedContent))
|
||||
}
|
||||
})
|
||||
|
||||
|
|
@ -158,10 +152,7 @@ func TestManagerUpload(t *testing.T) {
|
|||
fmt.Sprintf("bytes %d-%d/%d", start, end, totalSize))
|
||||
fileHeader.Header.Set("Content-Uid", "unique-file-id-123")
|
||||
|
||||
option := UploadOption{
|
||||
UserID: "user123",
|
||||
}
|
||||
|
||||
option := UploadOption{UserID: "user123"}
|
||||
file, err := manager.Upload(context.Background(), fileHeader, bytes.NewReader(chunk), option)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to upload chunk starting at %d: %v", start, err)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ package s3
|
|||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/gzip"
|
||||
"context"
|
||||
"fmt"
|
||||
"image"
|
||||
|
|
@ -229,6 +230,15 @@ func (storage *Storage) Reader(ctx context.Context, fileID string) (io.ReadClose
|
|||
return nil, fmt.Errorf("failed to get file: %w", err)
|
||||
}
|
||||
|
||||
// If the file is a gzip file, decompress it
|
||||
if strings.HasSuffix(fileID, ".gz") {
|
||||
reader, err := gzip.NewReader(result.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return reader, nil
|
||||
}
|
||||
|
||||
return result.Body, nil
|
||||
}
|
||||
|
||||
|
|
@ -254,6 +264,51 @@ func (storage *Storage) Download(ctx context.Context, fileID string) (io.ReadClo
|
|||
contentType = *result.ContentType
|
||||
}
|
||||
|
||||
// Try to detect content type from file extension
|
||||
ext := filepath.Ext(strings.TrimSuffix(fileID, ".gz"))
|
||||
switch strings.ToLower(ext) {
|
||||
case ".txt":
|
||||
contentType = "text/plain"
|
||||
case ".html":
|
||||
contentType = "text/html"
|
||||
case ".css":
|
||||
contentType = "text/css"
|
||||
case ".js":
|
||||
contentType = "application/javascript"
|
||||
case ".json":
|
||||
contentType = "application/json"
|
||||
case ".jpg", ".jpeg":
|
||||
contentType = "image/jpeg"
|
||||
case ".png":
|
||||
contentType = "image/png"
|
||||
case ".gif":
|
||||
contentType = "image/gif"
|
||||
case ".pdf":
|
||||
contentType = "application/pdf"
|
||||
case ".mp4":
|
||||
contentType = "video/mp4"
|
||||
case ".mp3":
|
||||
contentType = "audio/mpeg"
|
||||
case ".wav":
|
||||
contentType = "audio/wav"
|
||||
case ".ogg":
|
||||
contentType = "audio/ogg"
|
||||
case ".webm":
|
||||
contentType = "video/webm"
|
||||
case ".webp":
|
||||
contentType = "image/webp"
|
||||
case ".zip":
|
||||
}
|
||||
|
||||
// If the file is a gzip file, decompress it
|
||||
if strings.HasSuffix(fileID, ".gz") {
|
||||
reader, err := gzip.NewReader(result.Body)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
return reader, contentType, nil
|
||||
}
|
||||
|
||||
return result.Body, contentType, nil
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue