From 8bdafd115cee126a951587179422d8167fe6aa14 Mon Sep 17 00:00:00 2001 From: krau <71133316+krau@users.noreply.github.com> Date: Tue, 25 Aug 2026 15:22:59 +0800 Subject: [PATCH] fix(tdler): validate resume state against the part file --- common/tdler/resume.go | 21 +++++- common/tdler/resume_test.go | 111 ++++++++++++++++++++++++++++++ core/tasks/batchtfile/download.go | 14 ++-- core/tasks/batchtfile/execute.go | 14 ++-- core/tasks/tfile/download.go | 16 +++-- 5 files changed, 154 insertions(+), 22 deletions(-) diff --git a/common/tdler/resume.go b/common/tdler/resume.go index dfc2cdf..83da647 100644 --- a/common/tdler/resume.go +++ b/common/tdler/resume.go @@ -8,6 +8,7 @@ import ( "io" "net" "os" + "strings" "sync" "github.com/gotd/td/tg" @@ -84,7 +85,14 @@ func loadResumeBitmap(path string) (*resumeBitmap, error) { } var b resumeBitmap if err := json.Unmarshal(data, &b); err != nil { - return nil, fmt.Errorf("parse resume bitmap: %w", err) + // 无法解析的位图 (外部损坏): 删除并视为不存在, 全量重下自愈。 + _ = os.Remove(path) + return nil, nil + } + if b.Size <= 0 || b.PartSize <= 0 { + // 无效位图 (损坏或旧格式), 视为不存在, 全量重下。 + _ = os.Remove(path) + return nil, nil } b.ensureBlocks() return &b, nil @@ -182,6 +190,17 @@ func DownloadResumable( if err != nil { return err } + // 位图描述的数据文件 (bitmapPath 去掉 .bitmap 后缀) 必须存在且非空: + // 若缺失或为空, 已标记完成的块字节已丢失, 必须重置位图全量重下。 + if bm != nil { + partPath := strings.TrimSuffix(bitmapPath, ".bitmap") + if stat, err := os.Stat(partPath); err != nil || stat.Size() == 0 { + if err := os.Remove(bitmapPath); err != nil && !errors.Is(err, os.ErrNotExist) { + return fmt.Errorf("reset stale resume bitmap: %w", err) + } + bm = nil + } + } if bm == nil || bm.PartSize != tglimit.MaxPartSize || bm.Size != file.Size() { bm = newResumeBitmap(file.Size()) if err := bm.save(bitmapPath); err != nil { diff --git a/common/tdler/resume_test.go b/common/tdler/resume_test.go index 5eeea5e..934373c 100644 --- a/common/tdler/resume_test.go +++ b/common/tdler/resume_test.go @@ -129,6 +129,117 @@ func TestDownloadResumableBitmapResetOnSizeChange(t *testing.T) { } } +// TestDownloadResumablePartMissingOrTruncated resets the bitmap: skipped +// blocks would otherwise be zero-filled (caller recreates the part file +// without its bytes), or the download would wedge forever on a stale +// complete bitmap. +func TestDownloadResumablePartMissingOrTruncated(t *testing.T) { + data := make([]byte, 5*1024*1024) + for i := range data { + data[i] = byte(i % 251) + } + dir := t.TempDir() + partPath := filepath.Join(dir, "test.bin.part") + bitmapPath := ResumeStatePath(partPath) + + tests := []struct { + name string + doneBlocks []int + createPart bool + truncate bool + }{ + {"part missing, partial bitmap", []int{0, 1, 2}, false, false}, + {"part empty, partial bitmap", []int{0, 1, 2}, true, true}, + {"part missing, complete bitmap", []int{0, 1, 2, 3, 4}, false, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + os.Remove(partPath) + os.Remove(bitmapPath) + bm := newResumeBitmap(int64(len(data))) + for _, block := range tt.doneBlocks { + bm.markDone(block) + } + if err := bm.save(bitmapPath); err != nil { + t.Fatal(err) + } + if tt.createPart { + // Simulate the caller re-creating the part file (truncating). + if err := os.WriteFile(partPath, nil, 0o644); err != nil { + t.Fatal(err) + } + if tt.truncate { + if err := os.WriteFile(partPath, make([]byte, 0), 0o644); err != nil { + t.Fatal(err) + } + } + } + + partFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644) + if err != nil { + t.Fatal(err) + } + defer partFile.Close() + client := &serverLikeClient{data: data} + file := tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, client, int64(len(data)), "test.bin") + if err := DownloadResumable(context.Background(), file, partFile, 1, bitmapPath); err != nil { + t.Fatalf("download failed: %v", err) + } + got := make([]byte, len(data)) + if _, err := partFile.ReadAt(got, 0); err != nil { + t.Fatal(err) + } + if !bytesEqual(got, data) { + t.Fatalf("downloaded data mismatch (blocks not reset)") + } + }) + } +} + +// TestDownloadResumableInvalidBitmap treats a corrupt bitmap as absent. +func TestDownloadResumableInvalidBitmap(t *testing.T) { + data := make([]byte, 1024*1024+7) + for i := range data { + data[i] = byte(i % 251) + } + dir := t.TempDir() + partPath := filepath.Join(dir, "test.bin.part") + bitmapPath := ResumeStatePath(partPath) + for _, content := range []string{ + `{"part_size":1048576,"size":-1,"blocks":[]}`, + `{"part_size":1048576,"size":9223372036854775807,"blocks":[]}`, + `not json`, + } { + os.Remove(partPath) + if err := os.WriteFile(bitmapPath, []byte(content), 0o644); err != nil { + t.Fatal(err) + } + partFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644) + if err != nil { + t.Fatal(err) + } + client := &serverLikeClient{data: data} + file := tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, client, int64(len(data)), "test.bin") + err = DownloadResumable(context.Background(), file, partFile, 1, bitmapPath) + partFile.Close() + if err != nil { + t.Fatalf("download with corrupt bitmap %q failed: %v", content, err) + } + got := make([]byte, len(data)) + f, err := os.Open(partPath) + if err != nil { + t.Fatal(err) + } + if _, err := f.ReadAt(got, 0); err != nil { + t.Fatal(err) + } + f.Close() + if !bytesEqual(got, data) { + t.Fatalf("downloaded data mismatch with corrupt bitmap %q", content) + } + } +} + func TestRemoveResumeState(t *testing.T) { dir := t.TempDir() path := filepath.Join(dir, "x.bitmap") diff --git a/core/tasks/batchtfile/download.go b/core/tasks/batchtfile/download.go index c15dbf9..521410e 100644 --- a/core/tasks/batchtfile/download.go +++ b/core/tasks/batchtfile/download.go @@ -21,9 +21,11 @@ import ( // upload). func (t *Task) downloadToCache(ctx context.Context, elem *TaskElement) error { logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", elem.File.Name())) - if stat, err := os.Stat(elem.localPath); err == nil && stat.Size() == elem.File.Size() { - logger.Info("Cache file already complete, skipping download") - return nil + if elem.File.Size() > 0 { + if stat, err := os.Stat(elem.localPath); err == nil && stat.Size() == elem.File.Size() { + logger.Info("Cache file already complete, skipping download") + return nil + } } onProgress := t.downloadCallback(ctx, elem) if elem.File.Size() <= 0 { @@ -40,7 +42,8 @@ func (t *Task) downloadToCache(ctx context.Context, elem *TaskElement) error { return nil } partPath := elem.localPath + ".part" - localFile, err := fsutil.CreateFile(partPath) + // 不截断已存在的 .part: 位图标记的已完成块依赖既有字节。 + localFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644) if err != nil { return fmt.Errorf("failed to create local file: %w", err) } @@ -67,8 +70,9 @@ func (t *Task) downloadToCache(ctx context.Context, elem *TaskElement) error { if err := os.Rename(partPath, elem.localPath); err != nil { return fmt.Errorf("failed to finalize download: %w", err) } + // 清理位图是尽力而为: 下载已完成, 清理失败不应使任务失败。 if err := tdler.RemoveResumeState(tdler.ResumeStatePath(partPath)); err != nil { - return fmt.Errorf("failed to remove resume state: %w", err) + logger.Warnf("Failed to remove resume state: %v", err) } return nil } diff --git a/core/tasks/batchtfile/execute.go b/core/tasks/batchtfile/execute.go index 7b298a0..4b027cf 100644 --- a/core/tasks/batchtfile/execute.go +++ b/core/tasks/batchtfile/execute.go @@ -366,17 +366,12 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error { return nil } logger.Info("Starting file download") - localFile, err := fsutil.CreateFile(elem.localPath) - if err != nil { - t.markItemFailed(elem.ID, FailureStageCache, err) - t.notifyStateChange(ctx) - return fmt.Errorf("failed to create local file: %w", err) - } + // 不预创建缓存文件: 预创建会截断上次运行保留的完整缓存, 使复用失效。 success := false defer func() { if success { - if err := localFile.CloseAndRemove(); err != nil { - logger.Errorf("Failed to close local file: %v", err) + if err := os.Remove(elem.localPath); err != nil { + logger.Errorf("Failed to remove cache file: %v", err) } } }() @@ -392,8 +387,7 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error { elem.Path = elem.Path + ext } } - var fileStat os.FileInfo - fileStat, err = os.Stat(elem.localPath) + fileStat, err := os.Stat(elem.localPath) if err != nil { t.markItemFailed(elem.ID, FailureStageCache, err) t.notifyStateChange(ctx) diff --git a/core/tasks/tfile/download.go b/core/tasks/tfile/download.go index 66360b6..4532187 100644 --- a/core/tasks/tfile/download.go +++ b/core/tasks/tfile/download.go @@ -3,9 +3,9 @@ package tfile import ( "context" "fmt" - "github.com/charmbracelet/log" "os" + "github.com/charmbracelet/log" "github.com/krau/SaveAny-Bot/common/tdler" "github.com/krau/SaveAny-Bot/common/utils/dlutil" "github.com/krau/SaveAny-Bot/common/utils/fsutil" @@ -17,9 +17,11 @@ import ( // file (e.g. when the previous run was interrupted during upload). func (t *Task) download(ctx context.Context) error { logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", t.File.Name())) - if stat, err := os.Stat(t.localPath); err == nil && stat.Size() == t.File.Size() { - logger.Info("Cache file already complete, skipping download") - return nil + if t.File.Size() > 0 { + if stat, err := os.Stat(t.localPath); err == nil && stat.Size() == t.File.Size() { + logger.Info("Cache file already complete, skipping download") + return nil + } } if t.File.Size() <= 0 { // Unknown size (e.g. photos) cannot be resumed; use the plain downloader. @@ -36,7 +38,8 @@ func (t *Task) download(ctx context.Context) error { return nil } partPath := t.localPath + ".part" - localFile, err := fsutil.CreateFile(partPath) + // 不截断已存在的 .part: 位图标记的已完成块依赖既有字节。 + localFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644) if err != nil { return fmt.Errorf("failed to create local file: %w", err) } @@ -63,8 +66,9 @@ func (t *Task) download(ctx context.Context) error { if err := os.Rename(partPath, t.localPath); err != nil { return fmt.Errorf("failed to finalize download: %w", err) } + // 清理位图是尽力而为: 下载已完成, 清理失败不应使任务失败。 if err := tdler.RemoveResumeState(tdler.ResumeStatePath(partPath)); err != nil { - return fmt.Errorf("failed to remove resume state: %w", err) + logger.Warnf("Failed to remove resume state: %v", err) } logger.Info("File downloaded successfully") return nil