mirror of
https://github.com/krau/SaveAny-Bot.git
synced 2026-09-05 23:56:50 +08:00
fix(tdler): validate resume state against the part file
This commit is contained in:
+20
-1
@@ -8,6 +8,7 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"os"
|
"os"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/gotd/td/tg"
|
"github.com/gotd/td/tg"
|
||||||
@@ -84,7 +85,14 @@ func loadResumeBitmap(path string) (*resumeBitmap, error) {
|
|||||||
}
|
}
|
||||||
var b resumeBitmap
|
var b resumeBitmap
|
||||||
if err := json.Unmarshal(data, &b); err != nil {
|
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()
|
b.ensureBlocks()
|
||||||
return &b, nil
|
return &b, nil
|
||||||
@@ -182,6 +190,17 @@ func DownloadResumable(
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
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() {
|
if bm == nil || bm.PartSize != tglimit.MaxPartSize || bm.Size != file.Size() {
|
||||||
bm = newResumeBitmap(file.Size())
|
bm = newResumeBitmap(file.Size())
|
||||||
if err := bm.save(bitmapPath); err != nil {
|
if err := bm.save(bitmapPath); err != nil {
|
||||||
|
|||||||
@@ -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) {
|
func TestRemoveResumeState(t *testing.T) {
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
path := filepath.Join(dir, "x.bitmap")
|
path := filepath.Join(dir, "x.bitmap")
|
||||||
|
|||||||
@@ -21,10 +21,12 @@ import (
|
|||||||
// upload).
|
// upload).
|
||||||
func (t *Task) downloadToCache(ctx context.Context, elem *TaskElement) error {
|
func (t *Task) downloadToCache(ctx context.Context, elem *TaskElement) error {
|
||||||
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", elem.File.Name()))
|
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", elem.File.Name()))
|
||||||
|
if elem.File.Size() > 0 {
|
||||||
if stat, err := os.Stat(elem.localPath); err == nil && stat.Size() == elem.File.Size() {
|
if stat, err := os.Stat(elem.localPath); err == nil && stat.Size() == elem.File.Size() {
|
||||||
logger.Info("Cache file already complete, skipping download")
|
logger.Info("Cache file already complete, skipping download")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
}
|
||||||
onProgress := t.downloadCallback(ctx, elem)
|
onProgress := t.downloadCallback(ctx, elem)
|
||||||
if elem.File.Size() <= 0 {
|
if elem.File.Size() <= 0 {
|
||||||
// Unknown size (e.g. photos) cannot be resumed; use the plain downloader.
|
// Unknown size (e.g. photos) cannot be resumed; use the plain downloader.
|
||||||
@@ -40,7 +42,8 @@ func (t *Task) downloadToCache(ctx context.Context, elem *TaskElement) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
partPath := elem.localPath + ".part"
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create local file: %w", err)
|
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 {
|
if err := os.Rename(partPath, elem.localPath); err != nil {
|
||||||
return fmt.Errorf("failed to finalize download: %w", err)
|
return fmt.Errorf("failed to finalize download: %w", err)
|
||||||
}
|
}
|
||||||
|
// 清理位图是尽力而为: 下载已完成, 清理失败不应使任务失败。
|
||||||
if err := tdler.RemoveResumeState(tdler.ResumeStatePath(partPath)); err != nil {
|
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
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -366,17 +366,12 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
logger.Info("Starting file download")
|
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
|
success := false
|
||||||
defer func() {
|
defer func() {
|
||||||
if success {
|
if success {
|
||||||
if err := localFile.CloseAndRemove(); err != nil {
|
if err := os.Remove(elem.localPath); err != nil {
|
||||||
logger.Errorf("Failed to close local file: %v", err)
|
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
|
elem.Path = elem.Path + ext
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
var fileStat os.FileInfo
|
fileStat, err := os.Stat(elem.localPath)
|
||||||
fileStat, err = os.Stat(elem.localPath)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.markItemFailed(elem.ID, FailureStageCache, err)
|
t.markItemFailed(elem.ID, FailureStageCache, err)
|
||||||
t.notifyStateChange(ctx)
|
t.notifyStateChange(ctx)
|
||||||
|
|||||||
@@ -3,9 +3,9 @@ package tfile
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"github.com/charmbracelet/log"
|
|
||||||
"os"
|
"os"
|
||||||
|
|
||||||
|
"github.com/charmbracelet/log"
|
||||||
"github.com/krau/SaveAny-Bot/common/tdler"
|
"github.com/krau/SaveAny-Bot/common/tdler"
|
||||||
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
|
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
|
||||||
"github.com/krau/SaveAny-Bot/common/utils/fsutil"
|
"github.com/krau/SaveAny-Bot/common/utils/fsutil"
|
||||||
@@ -17,10 +17,12 @@ import (
|
|||||||
// file (e.g. when the previous run was interrupted during upload).
|
// file (e.g. when the previous run was interrupted during upload).
|
||||||
func (t *Task) download(ctx context.Context) error {
|
func (t *Task) download(ctx context.Context) error {
|
||||||
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", t.File.Name()))
|
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", t.File.Name()))
|
||||||
|
if t.File.Size() > 0 {
|
||||||
if stat, err := os.Stat(t.localPath); err == nil && stat.Size() == t.File.Size() {
|
if stat, err := os.Stat(t.localPath); err == nil && stat.Size() == t.File.Size() {
|
||||||
logger.Info("Cache file already complete, skipping download")
|
logger.Info("Cache file already complete, skipping download")
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
}
|
||||||
if t.File.Size() <= 0 {
|
if t.File.Size() <= 0 {
|
||||||
// Unknown size (e.g. photos) cannot be resumed; use the plain downloader.
|
// Unknown size (e.g. photos) cannot be resumed; use the plain downloader.
|
||||||
localFile, err := fsutil.CreateFile(t.localPath)
|
localFile, err := fsutil.CreateFile(t.localPath)
|
||||||
@@ -36,7 +38,8 @@ func (t *Task) download(ctx context.Context) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
partPath := t.localPath + ".part"
|
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 {
|
if err != nil {
|
||||||
return fmt.Errorf("failed to create local file: %w", err)
|
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 {
|
if err := os.Rename(partPath, t.localPath); err != nil {
|
||||||
return fmt.Errorf("failed to finalize download: %w", err)
|
return fmt.Errorf("failed to finalize download: %w", err)
|
||||||
}
|
}
|
||||||
|
// 清理位图是尽力而为: 下载已完成, 清理失败不应使任务失败。
|
||||||
if err := tdler.RemoveResumeState(tdler.ResumeStatePath(partPath)); err != nil {
|
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")
|
logger.Info("File downloaded successfully")
|
||||||
return nil
|
return nil
|
||||||
|
|||||||
Reference in New Issue
Block a user