diff --git a/chunk.go b/chunk.go new file mode 100644 index 0000000..65f2dba --- /dev/null +++ b/chunk.go @@ -0,0 +1,366 @@ +package main + +import ( + "encoding/json" + "io" + "log" + "net/http" + "os" + "path/filepath" + "regexp" + "strconv" + "sync" + "time" + + "github.com/gin-gonic/gin" +) + +const ( + chunkDir = "./data/chunks" + chunkThreshold = 5 << 30 // >5GB 才走分片 + chunkSize = 1 << 30 // 每片 1GB +) + +var uploadIDRe = regexp.MustCompile(`^[a-zA-Z0-9_-]{8,64}$`) + +type chunkMeta struct { + UploadID string `json:"upload_id"` + Filename string `json:"filename"` + TotalSize int64 `json:"total_size"` + TotalChunks int `json:"total_chunks"` + Received map[int]bool `json:"received"` + CreatedAt time.Time `json:"created_at"` + Reserved int64 `json:"reserved"` +} + +// 全局 map 存 in-flight 上传; 持锁只做元数据操作, 1GB 分片 io.Copy 在锁外进行. +// ponytail: 全局锁, 若并发大文件上传成为瓶颈再换 per-upload 锁. +var ( + chunkMu sync.RWMutex + chunkMetaMap = map[string]*chunkMeta{} +) + +func chunkDirPath(id string) string { return filepath.Join(chunkDir, id) } +func chunkFilePath(id string, idx int) string { + return filepath.Join(chunkDir, id, strconv.Itoa(idx)) +} + +func metaPath(id string) string { return filepath.Join(chunkDir, id, "meta.json") } + +// saveMeta 原子写 meta.json (.tmp + rename). +func saveMeta(m *chunkMeta) error { + b, err := json.Marshal(m) + if err != nil { + return err + } + tmp := metaPath(m.UploadID) + ".tmp" + if err := os.WriteFile(tmp, b, 0o644); err != nil { + return err + } + return os.Rename(tmp, metaPath(m.UploadID)) +} + +func loadMeta(id string) *chunkMeta { + b, err := os.ReadFile(metaPath(id)) + if err != nil { + return nil + } + var m chunkMeta + if err := json.Unmarshal(b, &m); err != nil { + return nil + } + return &m +} + +// expectedChunks totalSize 按 chunkSize 分片应有的片数 +func expectedChunks(totalSize int64) int { + return int((totalSize + chunkSize - 1) / chunkSize) +} + +// expectedChunkSize 第 i 片应有的字节数 (最后一片是余数) +func expectedChunkSize(totalSize int64, i int) int64 { + if n := totalSize - int64(i)*chunkSize; n < chunkSize { + return n + } + return chunkSize +} + +func uploadChunkHandler(c *gin.Context) { + uploadID := c.GetHeader("X-Upload-Id") + if !uploadIDRe.MatchString(uploadID) { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid upload id"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "invalid_upload_id"}) + return + } + idx, err := strconv.Atoi(c.GetHeader("X-Chunk-Index")) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid chunk index"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "invalid_chunk_index"}) + return + } + totalChunks, err := strconv.Atoi(c.GetHeader("X-Total-Chunks")) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid total chunks"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "invalid_total_chunks"}) + return + } + filename := c.GetHeader("X-Filename") + if !checkFilenameSafe(filename) { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid filename"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "invalid_filename"}) + return + } + totalSize, err := strconv.ParseInt(c.GetHeader("X-Total-Size"), 10, 64) + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "invalid total size"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "invalid_total_size"}) + return + } + if totalSize <= chunkThreshold { + c.JSON(http.StatusBadRequest, gin.H{"error": "chunked upload is for files > 5GB"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "below_threshold"}) + return + } + if totalChunks < 1 || idx < 0 || idx >= totalChunks { + c.JSON(http.StatusBadRequest, gin.H{"error": "chunk index out of range"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "index": idx, "total": totalChunks, "result": "index_out_of_range"}) + return + } + if e := expectedChunks(totalSize); totalChunks < e-1 || totalChunks > e+1 { + c.JSON(http.StatusBadRequest, gin.H{"error": "total chunks mismatch"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "total": totalChunks, "expected": e, "result": "chunk_count_mismatch"}) + return + } + if totalSize > quotaRemaining() { + c.JSON(http.StatusInsufficientStorage, gin.H{"error": "quota exceeded"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "size": totalSize, "result": "quota_exceeded"}) + return + } + + chunk, err := c.FormFile("chunk") + if err != nil { + c.JSON(http.StatusBadRequest, gin.H{"error": "missing form field 'chunk'"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "missing_form_field"}) + return + } + + dir := chunkDirPath(uploadID) + chunkMu.RLock() + meta, inMap := chunkMetaMap[uploadID] + chunkMu.RUnlock() + if !inMap { + meta = loadMeta(uploadID) // 进程重启后从磁盘恢复 + if meta != nil { + chunkMu.Lock() + chunkMetaMap[uploadID] = meta + chunkMu.Unlock() + } + } + if meta == nil { + if idx != 0 { + c.JSON(http.StatusNotFound, gin.H{"error": "upload not found"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "upload_not_found"}) + return + } + if err := os.MkdirAll(dir, 0o755); err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "mkdir_failed"}) + return + } + if !quotaReserve(totalSize) { + _ = os.RemoveAll(dir) + c.JSON(http.StatusInsufficientStorage, gin.H{"error": "quota exceeded"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "size": totalSize, "result": "quota_exceeded"}) + return + } + meta = &chunkMeta{ + UploadID: uploadID, Filename: filename, TotalSize: totalSize, + TotalChunks: totalChunks, Received: map[int]bool{}, + CreatedAt: time.Now(), Reserved: totalSize, + } + chunkMu.Lock() + chunkMetaMap[uploadID] = meta + err = saveMeta(meta) + chunkMu.Unlock() + if err != nil { + abortChunkUpload(uploadID, meta) + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "meta_save_failed"}) + return + } + } else if idx > 0 { + if _, err := os.Stat(dir); err != nil { + c.JSON(http.StatusNotFound, gin.H{"error": "upload not found"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "result": "upload_not_found"}) + return + } + } + + // 写分片: multipart → 流式 io.Copy 落盘, 不进内存 + dst := chunkFilePath(uploadID, idx) + in, err := chunk.Open() + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "index": idx, "result": "open_failed"}) + return + } + out, err := os.Create(dst) + if err != nil { + in.Close() + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "index": idx, "result": "create_failed"}) + return + } + _, cpErr := io.Copy(out, in) + in.Close() + if cerr := out.Close(); cerr != nil && cpErr == nil { + cpErr = cerr + } + if cpErr != nil { + _ = os.Remove(dst) + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "index": idx, "result": "write_failed"}) + return + } + + // 更新 Received + 持久化 (持写锁, 防并发分片时 meta.json 丢进度) + chunkMu.Lock() + meta.Received[idx] = true + err = saveMeta(meta) + complete := len(meta.Received) == meta.TotalChunks + chunkMu.Unlock() + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "index": idx, "result": "meta_save_failed"}) + return + } + + if !complete { + c.JSON(http.StatusAccepted, gin.H{"received": len(meta.Received), "total": meta.TotalChunks}) + audit(c, "upload_chunk", gin.H{"upload_id": uploadID, "index": idx, "received": len(meta.Received), "total": meta.TotalChunks, "result": "ok"}) + return + } + assembleChunks(c, uploadID, meta) +} + +// assembleChunks 所有分片到位后顺序拼装落盘并清理. +// ponytail: 末片重传/并发组装不额外去重, 前端单分片串行重试不会触发. +func assembleChunks(c *gin.Context, uploadID string, meta *chunkMeta) { + dir := chunkDirPath(uploadID) + for i := 0; i < meta.TotalChunks; i++ { + info, err := os.Stat(chunkFilePath(uploadID, i)) + if err != nil || info.Size() != expectedChunkSize(meta.TotalSize, i) { + abortChunkUpload(uploadID, meta) + c.JSON(http.StatusBadRequest, gin.H{"error": "chunk size mismatch"}) + audit(c, "upload_assembled", gin.H{"upload_id": uploadID, "index": i, "result": "chunk_size_mismatch"}) + return + } + } + + name := uniqueFilename(*uploadDir, meta.Filename) + dst := filepath.Join(*uploadDir, name) + out, err := os.Create(dst) + if err != nil { + abortChunkUpload(uploadID, meta) + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_assembled", gin.H{"upload_id": uploadID, "result": "create_failed"}) + return + } + var realSize int64 + for i := 0; i < meta.TotalChunks; i++ { + in, err := os.Open(chunkFilePath(uploadID, i)) + if err != nil { + out.Close() + _ = os.Remove(dst) + abortChunkUpload(uploadID, meta) + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_assembled", gin.H{"upload_id": uploadID, "index": i, "result": "open_failed"}) + return + } + n, cErr := io.Copy(out, in) + in.Close() + if cErr != nil { + out.Close() + _ = os.Remove(dst) + abortChunkUpload(uploadID, meta) + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_assembled", gin.H{"upload_id": uploadID, "index": i, "result": "copy_failed"}) + return + } + realSize += n + } + if err := out.Close(); err != nil { + _ = os.Remove(dst) + abortChunkUpload(uploadID, meta) + c.JSON(http.StatusInternalServerError, gin.H{"error": "internal error"}) + audit(c, "upload_assembled", gin.H{"upload_id": uploadID, "result": "close_failed"}) + return + } + + // 配额: 释放预留, 按真实大小重新占用 + quotaSub(meta.Reserved) + quotaReserve(realSize) + chunkMu.Lock() + delete(chunkMetaMap, uploadID) + chunkMu.Unlock() + _ = os.RemoveAll(dir) + + now := time.Now() + _ = os.Chtimes(dst, now, now) + c.JSON(http.StatusOK, fileResponse{ + Filename: name, URL: "/download/" + name, Size: realSize, + UploadedAt: now, ExpiresAt: now.Add(*fileTTL), + }) + audit(c, "upload_assembled", gin.H{ + "upload_id": uploadID, "file": name, "size": realSize, "result": "ok", + }) +} + +// abortChunkUpload 失败时释放预留配额并清掉分片目录与内存记录 +func abortChunkUpload(uploadID string, meta *chunkMeta) { + quotaSub(meta.Reserved) + chunkMu.Lock() + delete(chunkMetaMap, uploadID) + chunkMu.Unlock() + _ = os.RemoveAll(chunkDirPath(uploadID)) +} + +// cleanupChunks 清理超过 maxAge 的孤儿分片目录并释放预留配额 +func cleanupChunks(maxAge time.Duration) { + entries, err := os.ReadDir(chunkDir) + if err != nil { + if !os.IsNotExist(err) { + log.Printf("[cleaner] read chunk dir failed: %v", err) + } + return + } + now := time.Now() + for _, e := range entries { + if !e.IsDir() { + continue + } + id := e.Name() + dir := chunkDirPath(id) + if !uploadIDRe.MatchString(id) { + _ = os.RemoveAll(dir) // 非法目录名 (非本服务生成), 直接清 + continue + } + meta := loadMeta(id) + if meta == nil { + // 无 meta 的孤儿目录: 无法确定 reserved, 删目录即可 (quota 由 cleaner 的 rescanQuota 校正) + _ = os.RemoveAll(dir) + chunkMu.Lock() + delete(chunkMetaMap, id) + chunkMu.Unlock() + continue + } + if now.Sub(meta.CreatedAt) > maxAge { + quotaSub(meta.Reserved) + chunkMu.Lock() + delete(chunkMetaMap, id) + chunkMu.Unlock() + _ = os.RemoveAll(dir) + log.Printf("[cleaner] removed orphan chunks: %s (age=%v)", id, now.Sub(meta.CreatedAt)) + } + } +} diff --git a/chunk_test.go b/chunk_test.go new file mode 100644 index 0000000..ca5e8cd --- /dev/null +++ b/chunk_test.go @@ -0,0 +1,84 @@ +package main + +import ( + "bytes" + "mime/multipart" + "net/http" + "net/http/httptest" + "strconv" + "testing" + + "github.com/gin-gonic/gin" +) + +// 分片数学: 片数 ceil 与每片期望大小 +func TestChunkMath(t *testing.T) { + if got := expectedChunks(5<<30 + 1); got != 6 { + t.Fatalf("expectedChunks(5GB+1) = %d, want 6", got) + } + if got := expectedChunks(5 << 30); got != 5 { + t.Fatalf("expectedChunks(5GB) = %d, want 5", got) + } + if got := expectedChunkSize(5<<30+1, 0); got != chunkSize { + t.Fatalf("chunk0 size = %d, want 1GB", got) + } + if got := expectedChunkSize(5<<30+1, 5); got != 1 { + t.Fatalf("last chunk size = %d, want 1", got) + } +} + +func chunkRequest(t *testing.T, uploadID string, idx, totalChunks int, filename string, totalSize int64, content []byte) *gin.Context { + t.Helper() + var buf bytes.Buffer + w := multipart.NewWriter(&buf) + fw, err := w.CreateFormFile("chunk", "chunk") + if err != nil { + t.Fatal(err) + } + fw.Write(content) + w.Close() + req := httptest.NewRequest("POST", "/upload-chunk", &buf) + req.Header.Set("Content-Type", w.FormDataContentType()) + req.Header.Set("X-Upload-Id", uploadID) + req.Header.Set("X-Chunk-Index", strconv.Itoa(idx)) + req.Header.Set("X-Total-Chunks", strconv.Itoa(totalChunks)) + req.Header.Set("X-Filename", filename) + req.Header.Set("X-Total-Size", strconv.FormatInt(totalSize, 10)) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = req + return c +} + +// 校验拒绝路径: 非法 id / 小于阈值 / 越界 / 片数不匹配 / 超配额, 均不落盘 +func TestChunkHandlerRejections(t *testing.T) { + *auditDir = "" // 测试不写审计日志 + oldQuota := *quota + *quota = 1 << 30 + defer func() { *quota = oldQuota }() + + cases := []struct { + name string + id string + idx int + total int + fn string + size int64 + status int + }{ + {"invalid upload id", "bad id!", 0, 6, "x.bin", chunkThreshold + 1, http.StatusBadRequest}, + {"below threshold", "uploadid123456", 0, 1, "x.bin", chunkThreshold, http.StatusBadRequest}, + {"index out of range", "uploadid123456", 6, 6, "x.bin", chunkThreshold + 1, http.StatusBadRequest}, + {"chunk count mismatch", "uploadid123456", 0, 2, "x.bin", chunkThreshold + 1, http.StatusBadRequest}, + {"quota exceeded", "uploadid123456", 0, 6, "x.bin", chunkThreshold + 1, http.StatusInsufficientStorage}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + c := chunkRequest(t, tc.id, tc.idx, tc.total, tc.fn, tc.size, []byte("data")) + uploadChunkHandler(c) + if c.Writer.Status() != tc.status { + t.Fatalf("status = %d, want %d", c.Writer.Status(), tc.status) + } + }) + } +} diff --git a/cleaner.go b/cleaner.go index b205674..695c1ba 100644 --- a/cleaner.go +++ b/cleaner.go @@ -51,6 +51,15 @@ func startCleaner() { log.Fatalf("add cron job failed: %v", err) } c.Start() + // 分片目录独立清理: 复用同一周期, 清掉孤儿分片并释放预留配额 + go func() { + cleanupChunks(*fileTTL) // 启动先清一轮, 处理上次进程遗留 + ticker := time.NewTicker(*scanEvery) + defer ticker.Stop() + for range ticker.C { + cleanupChunks(*fileTTL) + } + }() log.Printf("[cleaner] started, schedule=@every %s dir=%s ttl=%s", scanEvery, *uploadDir, *fileTTL) }) diff --git a/handlers.go b/handlers.go index 87ff228..16da746 100644 --- a/handlers.go +++ b/handlers.go @@ -24,6 +24,12 @@ type fileResponse struct { ExpiresAt time.Time `json:"expires_at"` } +// checkFilenameSafe 校验文件名: 不含 NUL 字节, 且不超过 200 字节 (Linux NAME_MAX=255 的余量). +// 供 uploadHandler 与 uploadChunkHandler 复用 (NUL/200 校验只写这一份). +func checkFilenameSafe(name string) bool { + return !strings.ContainsRune(name, 0) && len(name) <= 200 +} + // validateFilename 从 URL/JSON 参数提取并校验 filename. // 返回 (sanitized, true) 表示安全可访问, 否则返回 ("", false). // 规则与 downloadHandler 原内联逻辑一致: filepath.Base + 拒绝 "" / "." / ".." @@ -44,14 +50,14 @@ func uploadHandler(c *gin.Context) { audit(c, "upload", gin.H{"file": "", "result": "missing_form_field"}) return } - // 早 reject: 防止 NUL 字节 + 超长文件名打爆 Stat/Save - if strings.ContainsRune(f.Filename, 0) { - c.JSON(http.StatusBadRequest, gin.H{"error": "filename contains null byte"}) - audit(c, "upload", gin.H{"file": f.Filename, "result": "null_byte"}) - return - } - // Linux/macOS 路径分量 NAME_MAX = 255 字节 - if len(f.Filename) > 200 { + // 早 reject: 防止 NUL 字节 + 超长文件名打爆 Stat/Save (规则与分片上传共用) + if !checkFilenameSafe(f.Filename) { + if strings.ContainsRune(f.Filename, 0) { + c.JSON(http.StatusBadRequest, gin.H{"error": "filename contains null byte"}) + audit(c, "upload", gin.H{"file": f.Filename, "result": "null_byte"}) + return + } + // Linux/macOS 路径分量 NAME_MAX = 255 字节 c.JSON(http.StatusBadRequest, gin.H{"error": "filename too long (max 200 bytes)"}) audit(c, "upload", gin.H{"file": f.Filename, "result": "filename_too_long"}) return diff --git a/main.go b/main.go index f0bcbd3..d9ccf58 100644 --- a/main.go +++ b/main.go @@ -58,6 +58,7 @@ func main() { // 业务路由 r.POST("/upload", uploadHandler) + r.POST("/upload-chunk", uploadChunkHandler) r.GET("/download/:filename", downloadHandler) r.POST("/download-zip", batchDownloadHandler) r.POST("/files-delete", batchDeleteHandler) diff --git a/static/app.js b/static/app.js index ab782d7..91763ff 100644 --- a/static/app.js +++ b/static/app.js @@ -90,6 +90,8 @@ function addFiles(fileList) { id: nextId++, file: f, xhr: null, + chunkXhrs: new Set(), + chunkProgress: null, status: 'pending', progress: 0, loaded: 0, @@ -117,9 +119,9 @@ function renderPending() { item.status === 'error' ? '✗' : item.status === 'cancelled' ? '⊘' : ''; - const sizeInfo = item.file.size > 0 + const sizeInfo = (item.file.size > 0 ? `${formatSize(item.loaded)} / ${formatSize(item.file.size)}` - : formatSize(item.file.size); + : formatSize(item.file.size)) + (item.file.size > CHUNK_THRESHOLD ? ' (分片上传)' : ''); div.innerHTML = `