feat(core): persist tasks and recover them after restart

This commit is contained in:
krau
2026-08-25 15:05:55 +08:00
parent 54dc4caafe
commit 991f454096
14 changed files with 938 additions and 68 deletions
+106
View File
@@ -0,0 +1,106 @@
package batchtfile
import (
"context"
"encoding/json"
"fmt"
"path/filepath"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/core"
"github.com/krau/SaveAny-Bot/pkg/enums/tasktype"
tfilepkg "github.com/krau/SaveAny-Bot/pkg/tfile"
"github.com/krau/SaveAny-Bot/storage"
)
type elementPayload struct {
ID string `json:"id"`
Storage string `json:"storage"`
Path string `json:"path"`
File tfilepkg.FilePayload `json:"file"`
SourceGroupKey string `json:"source_group_key"`
SourceCaption string `json:"source_caption"`
PreserveCaption bool `json:"preserve_caption"`
}
type taskPayload struct {
ID string `json:"id"`
Elements []elementPayload `json:"elements"`
ChatID int64 `json:"chat_id"`
MessageID int `json:"message_id"`
IgnoreErrors bool `json:"ignore_errors"`
}
type taskCodec struct{}
func init() {
core.RegisterTaskCodec(tasktype.TaskTypeTgfiles, taskCodec{})
}
func (taskCodec) Marshal(task core.Executable) ([]byte, error) {
t, ok := task.(*Task)
if !ok {
return nil, fmt.Errorf("unexpected task type %T", task)
}
p := taskPayload{
ID: t.ID,
IgnoreErrors: t.IgnoreErrors,
}
for _, elem := range t.elems {
filePayload, ok := tfilepkg.FilePayloadOf(elem.File)
if !ok {
return nil, fmt.Errorf("file %T is not serializable", elem.File)
}
p.Elements = append(p.Elements, elementPayload{
ID: elem.ID,
Storage: elem.Storage.Name(),
Path: elem.Path,
File: filePayload,
SourceGroupKey: elem.sourceGroupKey,
SourceCaption: elem.sourceCaption,
PreserveCaption: elem.preserveCaption,
})
}
if progress, ok := t.Progress.(*Progress); ok {
p.ChatID = progress.ChatID
p.MessageID = progress.MessageID
}
return json.Marshal(p)
}
func (taskCodec) Unmarshal(data []byte) (core.Executable, error) {
var p taskPayload
if err := json.Unmarshal(data, &p); err != nil {
return nil, fmt.Errorf("invalid task payload: %w", err)
}
dler := core.DownloaderClient()
if dler == nil {
return nil, fmt.Errorf("no downloader client available")
}
elems := make([]TaskElement, 0, len(p.Elements))
for _, ep := range p.Elements {
stor, err := storage.GetStorageByName(context.Background(), ep.Storage)
if err != nil {
return nil, fmt.Errorf("storage %q: %w", ep.Storage, err)
}
localPath, err := filepath.Abs(filepath.Join(config.C().Temp.BasePath, fmt.Sprintf("%s_%s", ep.ID, ep.File.Name)))
if err != nil {
return nil, fmt.Errorf("failed to build cache path: %w", err)
}
elems = append(elems, TaskElement{
ID: ep.ID,
Storage: stor,
Path: ep.Path,
File: tfilepkg.FileFromPayload(ep.File, dler),
localPath: localPath,
sourceGroupKey: ep.SourceGroupKey,
sourceCaption: ep.SourceCaption,
preserveCaption: ep.PreserveCaption,
})
}
var progress ProgressTracker
if p.ChatID != 0 {
progress = NewProgressTracker(p.MessageID, p.ChatID)
}
return NewBatchTGFileTask(p.ID, context.Background(), elems, progress, p.IgnoreErrors), nil
}
+20 -43
View File
@@ -134,7 +134,12 @@ func (t *Task) processElements(ctx context.Context, elems []*TaskElement) error
}
func (t *Task) processBatch(ctx context.Context, group executionGroup) error {
// Cache files are kept on failure so a later restart can resume upload.
uploaded := false
defer func() {
if !uploaded {
return
}
for _, elem := range group.elems {
if err := os.Remove(elem.localPath); err != nil && !os.IsNotExist(err) {
log.FromContext(ctx).Warnf("Failed to cleanup batch cache file %s: %v", elem.localPath, err)
@@ -219,7 +224,11 @@ func (t *Task) processBatch(ctx context.Context, group executionGroup) error {
for index, item := range items {
t.recordDownloadComplete(successElems[index].ID, item.Size)
}
return t.saveBatchItems(ctx, successElems, items)
err := t.saveBatchItems(ctx, successElems, items)
if err == nil {
uploaded = true
}
return err
}
func (t *Task) saveBatchItems(ctx context.Context, successElems []*TaskElement, items []storagetypes.BatchItem) error {
@@ -289,34 +298,10 @@ func (t *Task) unmarkProcessing(id string) {
func (t *Task) downloadElement(ctx context.Context, elem *TaskElement) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", elem.File.Name()))
logger.Info("Starting file download")
localFile, err := fsutil.CreateFile(elem.localPath)
if err != nil {
t.markItemFailed(elem.ID, FailureStageCache, err)
if err := t.downloadToCache(ctx, elem); err != nil {
t.markItemFailed(elem.ID, FailureStageDownload, err)
t.notifyStateChange(ctx)
return fmt.Errorf("failed to create local file: %w", err)
}
wrAt := ioutil.NewProgressWriterAt(localFile, func(n int) {
t.recordItemDownload(elem.ID, int64(n), time.Now())
downloaded := t.downloaded.Add(int64(n))
t.notifyProgress(ctx)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
TotalBytes: t.totalSize,
DownloadedBytes: downloaded,
})
})
_, downloadErr := tdler.NewDownloader(elem.File).Parallel(ctx, wrAt)
closeErr := localFile.Close()
if downloadErr != nil {
t.markItemFailed(elem.ID, FailureStageDownload, downloadErr)
t.notifyStateChange(ctx)
return fmt.Errorf("failed to download file: %w", downloadErr)
}
if closeErr != nil {
t.markItemFailed(elem.ID, FailureStageCache, closeErr)
t.notifyStateChange(ctx)
return fmt.Errorf("failed to close cache file: %w", closeErr)
return fmt.Errorf("failed to download file: %w", err)
}
logger.Info("File downloaded successfully")
if path.Ext(elem.FileName()) == "" {
@@ -387,24 +372,15 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
t.notifyStateChange(ctx)
return fmt.Errorf("failed to create local file: %w", err)
}
success := false
defer func() {
if err := localFile.CloseAndRemove(); err != nil {
logger.Errorf("Failed to close local file: %v", err)
if success {
if err := localFile.CloseAndRemove(); err != nil {
logger.Errorf("Failed to close local file: %v", err)
}
}
}()
wrAt := ioutil.NewProgressWriterAt(localFile, func(n int) {
t.recordItemDownload(elem.ID, int64(n), time.Now())
downloaded := t.downloaded.Add(int64(n))
t.notifyProgress(ctx)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
TotalBytes: t.totalSize,
DownloadedBytes: downloaded,
})
})
_, err = tdler.NewDownloader(elem.File).Parallel(ctx, wrAt)
if err != nil {
if err := t.downloadToCache(ctx, &elem); err != nil {
t.markItemFailed(elem.ID, FailureStageDownload, err)
t.notifyStateChange(ctx)
return fmt.Errorf("failed to download file: %w", err)
@@ -460,6 +436,7 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
onProgress(fileStat.Size(), fileStat.Size())
t.markItemCompleted(elem.ID)
t.notifyStateChange(vctx)
success = true
} else {
t.markItemFailed(elem.ID, lastFailureStage, err)
t.notifyStateChange(vctx)