From c2f8ab3c0118074e382e40222213b047689ba798 Mon Sep 17 00:00:00 2001
From: Krau <71133316+krau@users.noreply.github.com>
Date: Mon, 17 Aug 2026 19:10:03 +0800
Subject: [PATCH] refactor: quality overhaul (#234)
* fix: prevent queue deadlock after cancelling tasks
Get no longer recurses while holding the mutex.
Cancelled queued tasks leave the map, making their IDs reusable.
Closed empty queues return ErrQueueClosed.
* fix: init task queue lazily and close on shutdown
AddTask is safe before Run by initializing the queue once.
Close unblocks workers waiting in Get.
* fix: make resource fingerprints deterministic
Sort map keys before hashing so Resource.ID is stable.
* fix: harden per-item processing tracking in tasks
Use the resource fingerprint as the dedup key on insert and delete.
Check and set processing entries atomically instead of TOCTOU.
Count only successful downloads.
* fix: honor IgnoreErrors in batch tasks
Element failures no longer cancel sibling elements or stop later groups.
* fix: report streamed upload bytes in batch tasks
Stream downloads report their byte count as the upload total.
* fix: validate i18n key parity across locales
geni18n fails when a language file misses any key.
Align the syncpeers completion key between en and zh-Hans.
Translate three untranslated parse keys in zh-Hans.
* refactor: share progress throttling helpers
Move size-tiered and count-based throttling into progressutil.
Localize hardcoded Chinese progress strings.
Drop five duplicated implementations and dead local copies.
* refactor: share unique filename logic
Storage backends use fsutil.UniquePath instead of local loops.
* fix: sanitize local storage paths
Reject absolute paths and dot-dot escapes in Save.
Check close errors and wrap creation failures.
* fix: preserve webdav error causes
Wrap mkdir and write failures with %w and drop dead error values.
* fix: kill rclone subprocess on reader close
Prevent cat processes from hanging after the pipe is closed.
* fix: fail alist init instead of exiting
Replace log.Fatalf with wrapped errors so a bad alist cannot kill the bot.
Cancel token refresh with the init context and re-login on 401.
* fix: enforce telegram album limits and track saved paths
Use tglimit.MaxAlbumItems for album batching and video splitting.
Exists now reports previously saved paths instead of always false.
* fix: guard storage registry maps
Protect Storages and UserStorages with mutexes and expose read accessors.
* fix: harden JS parser plugin runtime
Validate semver instead of panicking and require canHandle.
Recover plugin worker panics and time out CanHandle calls.
Return a copy from the registry and sanitize install filenames.
* fix: guard kemono parser against nil fields
Skip sparse preview and attachment entries instead of panicking.
* refactor: remove commented-out dead code
Drop unused kemono legacy types and commented response structs.
* test: cover invalid plugin version rejection
* fix: show all queued tasks in /task list
Render up to ten tasks and append the truncation note once.
* fix: guard callback data parsing
Reject malformed callback payloads before indexing split parts.
* fix: isolate media groups per user
Key pending groups by chat, user and group id.
* fix: require permission for callback handlers
* fix: initialize userbot context once
Replace the racy lazy init with sync.OnceValue.
* fix: avoid leaking raw errors in /dir reply
* fix: notify users on invalid update version
* fix: fail fast when API listen fails
Bind synchronously and surface errors instead of logging them.
* fix: add timeouts and backoff to webhook delivery
* refactor: remove dead code from api and bot
Drop the empty ProgressTracker shim, unused token context key and a redundant SetBotCommands call.
* fix: load remote config without local lookup
Skip the local file search after reading a config URL and add a timeout.
* refactor: drop unused hook config
* docs: document parser plugin config
* ci: fix BuildTime formatting and align checkout
Actions format does not format dates; pass the raw timestamp.
* chore: ignore cache directory
* fix: make cache init idempotent
* fix: upload only downloaded batch elements
Failed elements keep partial cache files and never reach the backend.
Successfully downloaded siblings still upload when one element fails.
* fix: record telegram saved paths only after upload
Skip-large returns a sentinel so skipped files are not marked as saved.
* fix: deduplicate alist token refreshes
Guard token access with a mutex and merge concurrent logins.
Reuse a recent refresh to avoid login storms.
* test: cover concurrent alist 401 retry
Ten parallel uploads share a single re-login under -race.
* fix: count parsed resources in progress text
* fix: localize storage lookup errors in /dir
Use the shared i18n key and escape the dynamic error.
* fix: deduplicate concurrent storage initialization
singleflight merges first-time inits so side effects are not duplicated.
* fix: guard nil progress trackers in api-created tasks
Batch, telegraph and transfer tasks run without a Telegram tracker
when created through the API; their callbacks must not panic.
* fix: keep album order when filtering failed batch items
Download results are stored by original index so surviving
elements keep their source order.
* test: cover nil-tracker task execution
* fix: report partial failure in batch done message
IgnoreErrors runs with failed elements show success and failed
counts instead of claiming every file completed.
* fix: skip login refresh for token-only alist storage
401 responses surface the auth error instead of sending a
credential-less login; the refresh window never exceeds TokenExp.
* test: cover alist refresh semantics
Concurrent refresh uses username/password; token-only storage
never attempts a login on 401.
* fix: call tracker from notifyProgress instead of recursing
The helper called itself, overflowing the stack on any batch
task with a progress tracker.
* fix: propagate cancellation past IgnoreErrors
Cancelled tasks must not be reported as successful: only
ordinary element failures are ignored.
* test: cover notifyProgress tracker call
* test: cover cancellation with IgnoreErrors
* fix: refresh alist token after startup 401s
The init login no longer satisfies the dedup window, so an
early 401 triggers a real refresh; later 401s reuse it.
Clarify that streaming uploads cannot replay their body.
* test: exercise the alist 401 refresh path
Concurrent uploads now reject the init token, share one
refresh and retry with the new token; token-only stays inert.
* style: trim verbose comments
Drop process-style explanations; keep one-line behavior notes.
* style: gofmt test files
* fix: drop unbounded saved-path cache in telegram storage
Telegram cannot reliably query remote file existence, so the
cache answered a question it could not answer and grew without
bound. Exists returns false again, as before.
* style: format codes
---
.github/workflows/build-release.yml | 4 +-
.gitignore | 1 +
api/auth.go | 8 +-
api/factory.go | 6 +-
api/handlers.go | 5 +-
api/progress.go | 16 ---
api/server.go | 9 +-
api/webhook.go | 31 +++--
client/bot/bot.go | 3 -
client/bot/handlers/add_task.go | 6 +-
client/bot/handlers/cancel_task.go | 7 +-
client/bot/handlers/dir.go | 5 +-
client/bot/handlers/media_group.go | 45 ++++---
client/bot/handlers/middleware.go | 11 ++
client/bot/handlers/register.go | 10 +-
client/bot/handlers/silent.go | 7 +-
client/bot/handlers/tasks.go | 13 +-
client/bot/handlers/update.go | 6 +-
client/user/userclient.go | 14 ++-
cmd/geni18n/main.go | 32 ++++-
cmd/run.go | 3 +-
common/cache/ristretto.go | 33 ++---
common/i18n/i18nk/keys.go | 5 +-
common/i18n/locale/en.yaml | 4 +-
common/i18n/locale/zh-Hans.yaml | 8 +-
common/utils/fsutil/unique.go | 27 ++++
common/utils/progressutil/progress.go | 50 ++++++++
common/utils/tgutil/message.go | 91 --------------
config.example.toml | 19 +++
config/hook.go | 9 --
config/viper.go | 23 +++-
core/core.go | 29 ++++-
core/tasks/batchtfile/execute.go | 110 ++++++++++++----
core/tasks/batchtfile/notify_test.go | 35 ++++++
core/tasks/batchtfile/progress.go | 9 +-
core/tasks/batchtfile/task.go | 2 -
core/tasks/batchtfile/utils.go | 32 -----
core/tasks/directlinks/execute.go | 7 +-
core/tasks/directlinks/progress.go | 8 +-
core/tasks/directlinks/task.go | 2 -
core/tasks/directlinks/util.go | 31 -----
core/tasks/parsed/execute.go | 16 +--
core/tasks/parsed/progress.go | 44 ++-----
core/tasks/parsed/task.go | 2 -
core/tasks/telegraph/execute.go | 12 +-
core/tasks/telegraph/nil_tracker_test.go | 31 +++++
core/tasks/telegraph/progress.go | 3 +-
core/tasks/telegraph/utils.go | 13 --
core/tasks/tfile/progress.go | 5 +-
core/tasks/tfile/util.go | 32 -----
core/tasks/transfer/cancel_test.go | 75 +++++++++++
core/tasks/transfer/execute.go | 22 ++--
core/tasks/transfer/progress.go | 11 +-
parsers/js/api.go | 8 +-
parsers/js/js.go | 107 ++++++++++------
parsers/js/plugin_version_test.go | 35 ++++++
parsers/native/kemono/kemono.go | 9 ++
parsers/native/kemono/post_legacy.go | 16 ---
parsers/native/kemono/user_profile.go | 8 --
parsers/native/twitter/types.go | 96 --------------
parsers/parsers/parsers.go | 3 +-
pkg/consts/tglimit/tglimit.go | 3 +
pkg/parser/parser.go | 11 +-
pkg/queue/queue.go | 37 +++---
pkg/queue/queue_test.go | 86 ++++++++++++-
storage/alist/alist.go | 139 +++++++++++++--------
storage/alist/alist_test.go | 152 +++++++++++++++++++++++
storage/alist/token.go | 56 ++++++++-
storage/load.go | 79 ++++++++++--
storage/local/local.go | 27 ++--
storage/minio/client.go | 18 +--
storage/rclone/errs.go | 3 -
storage/rclone/rclone.go | 24 ++--
storage/s3/s3.go | 18 +--
storage/storage.go | 11 +-
storage/telegram/telegram.go | 13 +-
storage/telegram/video_split.go | 3 +-
storage/webdav/errs.go | 9 --
storage/webdav/webdav.go | 27 ++--
79 files changed, 1276 insertions(+), 764 deletions(-)
create mode 100644 common/utils/fsutil/unique.go
create mode 100644 common/utils/progressutil/progress.go
create mode 100644 core/tasks/batchtfile/notify_test.go
delete mode 100644 core/tasks/batchtfile/utils.go
create mode 100644 core/tasks/telegraph/nil_tracker_test.go
delete mode 100644 core/tasks/telegraph/utils.go
delete mode 100644 core/tasks/tfile/util.go
create mode 100644 core/tasks/transfer/cancel_test.go
create mode 100644 parsers/js/plugin_version_test.go
delete mode 100644 parsers/native/kemono/post_legacy.go
delete mode 100644 parsers/native/kemono/user_profile.go
create mode 100644 storage/alist/alist_test.go
delete mode 100644 storage/webdav/errs.go
diff --git a/.github/workflows/build-release.yml b/.github/workflows/build-release.yml
index 3d97791..480861c 100644
--- a/.github/workflows/build-release.yml
+++ b/.github/workflows/build-release.yml
@@ -43,7 +43,7 @@ jobs:
goarch: arm64
steps:
- name: Checkout
- uses: actions/checkout@v4
+ uses: actions/checkout@v6
- name: Extract version from Git Ref
id: extract_version
@@ -64,7 +64,7 @@ jobs:
ldflags: >-
-s -w
-X "github.com/krau/SaveAny-Bot/config.Version=${{ env.VERSION }}"
- -X "github.com/krau/SaveAny-Bot/config.BuildTime=${{ format(github.event.repository.updated_at, 'yyyy-MM-dd HH:mm:ss') }}"
+ -X "github.com/krau/SaveAny-Bot/config.BuildTime=${{ github.event.repository.updated_at }}"
-X "github.com/krau/SaveAny-Bot/config.GitCommit=${{ github.sha }}"
binary_name: saveany-bot
env:
diff --git a/.gitignore b/.gitignore
index 412da25..d4398b1 100644
--- a/.gitignore
+++ b/.gitignore
@@ -1,5 +1,6 @@
config.toml
logs/
+/cache/
tmp/
data/
downloads/
diff --git a/api/auth.go b/api/auth.go
index 236bec3..49d2179 100644
--- a/api/auth.go
+++ b/api/auth.go
@@ -1,7 +1,6 @@
package api
import (
- "context"
"crypto/subtle"
"net/http"
"strings"
@@ -9,9 +8,6 @@ import (
"github.com/krau/SaveAny-Bot/config"
)
-// tokenContextKey 用于在 context 中存储 token
-type tokenContextKey struct{}
-
// AuthMiddleware 返回认证中间件
func AuthMiddleware() func(http.Handler) http.Handler {
return func(next http.Handler) http.Handler {
@@ -40,9 +36,7 @@ func AuthMiddleware() func(http.Handler) http.Handler {
return
}
- // 将 token 添加到 context
- ctx := context.WithValue(r.Context(), tokenContextKey{}, token)
- next.ServeHTTP(w, r.WithContext(ctx))
+ next.ServeHTTP(w, r)
})
}
}
diff --git a/api/factory.go b/api/factory.go
index 85bb2f5..6c55a2d 100644
--- a/api/factory.go
+++ b/api/factory.go
@@ -39,7 +39,7 @@ func NewTaskFactory(ctx context.Context) *TaskFactory {
// CreateTask 创建任务
func (f *TaskFactory) CreateTask(req *CreateTaskRequest) (*CreateTaskResponse, error) {
// 验证存储
- stor, ok := storage.Storages[req.Storage]
+ stor, ok := storage.GetStorage(req.Storage)
if !ok {
return nil, fmt.Errorf("storage not found: %s", req.Storage)
}
@@ -327,12 +327,12 @@ func (f *TaskFactory) createTransferTask(taskID string, createdAt time.Time, req
}
// 验证源存储和目标存储
- sourceStor, ok := storage.Storages[params.SourceStorage]
+ sourceStor, ok := storage.GetStorage(params.SourceStorage)
if !ok {
return nil, fmt.Errorf("source storage not found: %s", params.SourceStorage)
}
- targetStor, ok := storage.Storages[params.TargetStorage]
+ targetStor, ok := storage.GetStorage(params.TargetStorage)
if !ok {
return nil, fmt.Errorf("target storage not found: %s", params.TargetStorage)
}
diff --git a/api/handlers.go b/api/handlers.go
index 5f4f83b..53b65a9 100644
--- a/api/handlers.go
+++ b/api/handlers.go
@@ -135,8 +135,9 @@ func (h *Handlers) ListStoragesHandler(w http.ResponseWriter, r *http.Request) {
return
}
- storages := make([]StorageInfo, 0, len(storage.Storages))
- for name, stor := range storage.Storages {
+ all := storage.AllStorages()
+ storages := make([]StorageInfo, 0, len(all))
+ for name, stor := range all {
storages = append(storages, StorageInfo{
Name: name,
Type: string(stor.Type()),
diff --git a/api/progress.go b/api/progress.go
index d592a84..e6e5a52 100644
--- a/api/progress.go
+++ b/api/progress.go
@@ -198,19 +198,3 @@ func (t *TaskProgressInfo) Emit(e taskevent.Event) {
SendWebhook(nil, payload)
}
}
-
-// ProgressTracker is retained for compatibility but is no longer the primary
-// progress path; taskevent drives updates now. These methods are safe no-ops
-// when called on a nil receiver.
-type ProgressTracker struct{}
-
-func NewProgressTracker(taskID, taskType, storage, path, title, webhook string) *ProgressTracker {
- return &ProgressTracker{}
-}
-
-func (p *ProgressTracker) OnStart(totalBytes int64, totalFiles int) {}
-func (p *ProgressTracker) OnProgress(downloadedBytes int64, downloadedFiles int) {}
-func (p *ProgressTracker) OnDone(err error) {}
-func (p *ProgressTracker) GetInfo() *TaskProgressInfo { return nil }
-func (p *ProgressTracker) UpdateProgressBytes(bytes int64) {}
-func (p *ProgressTracker) UpdateProgressFiles(files int) {}
diff --git a/api/server.go b/api/server.go
index 9418e1e..a0733dd 100644
--- a/api/server.go
+++ b/api/server.go
@@ -3,6 +3,7 @@ package api
import (
"context"
"fmt"
+ "net"
"net/http"
"time"
@@ -90,9 +91,15 @@ func (s *Server) Start(ctx context.Context) error {
logger.Infof("Starting API server on %s", s.httpServer.Addr)
+ // Bind synchronously so listen failures are returned to the caller.
+ ln, err := net.Listen("tcp", s.httpServer.Addr)
+ if err != nil {
+ return fmt.Errorf("failed to listen on %s: %w", s.httpServer.Addr, err)
+ }
+
// 在 goroutine 中启动服务器
go func() {
- if err := s.httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed {
+ if err := s.httpServer.Serve(ln); err != nil && err != http.ErrServerClosed {
logger.Errorf("API server error: %v", err)
}
}()
diff --git a/api/webhook.go b/api/webhook.go
index 7e4ad32..7733b6a 100644
--- a/api/webhook.go
+++ b/api/webhook.go
@@ -37,6 +37,9 @@ func SendWebhook(ctx context.Context, payload *WebhookPayload) {
} else {
logger = log.Default().With("task_id", payload.TaskID)
}
+ if ctx == nil {
+ ctx = context.Background()
+ }
payloadBytes, err := json.Marshal(payload)
if err != nil {
@@ -44,10 +47,15 @@ func SendWebhook(ctx context.Context, payload *WebhookPayload) {
return
}
- // 重试 3 次
- for i := range 3 {
- req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, webhookURL, bytes.NewBuffer(payloadBytes))
+ // 重试 3 次, 指数退避 (100ms/400ms/1.6s)
+ const maxAttempts = 3
+ const requestTimeout = 30 * time.Second
+ backoff := 100 * time.Millisecond
+ for i := range maxAttempts {
+ reqCtx, cancel := context.WithTimeout(ctx, requestTimeout)
+ req, err := http.NewRequestWithContext(reqCtx, http.MethodPost, webhookURL, bytes.NewBuffer(payloadBytes))
if err != nil {
+ cancel()
logger.Errorf("Failed to create webhook request: %v", err)
return
}
@@ -56,9 +64,13 @@ func SendWebhook(ctx context.Context, payload *WebhookPayload) {
req.Header.Set("User-Agent", "SaveAny-Bot/1.0")
resp, err := webhookClient.Do(req)
+ cancel()
if err != nil {
- logger.Warnf("Webhook request failed (attempt %d/3): %v", i+1, err)
- time.Sleep(time.Second * time.Duration(i+1))
+ logger.Warnf("Webhook request failed (attempt %d/%d): %v", i+1, maxAttempts, err)
+ if i < maxAttempts-1 {
+ time.Sleep(backoff)
+ }
+ backoff *= 4
continue
}
resp.Body.Close()
@@ -68,11 +80,14 @@ func SendWebhook(ctx context.Context, payload *WebhookPayload) {
return
}
- logger.Warnf("Webhook returned non-2xx status (attempt %d/3): %d", i+1, resp.StatusCode)
- time.Sleep(time.Second * time.Duration(i+1))
+ logger.Warnf("Webhook returned non-2xx status (attempt %d/%d): %d", i+1, maxAttempts, resp.StatusCode)
+ if i < maxAttempts-1 {
+ time.Sleep(backoff)
+ }
+ backoff *= 4
}
- logger.Errorf("Failed to send webhook after 3 attempts")
+ logger.Errorf("Failed to send webhook after %d attempts", maxAttempts)
}()
}
diff --git a/client/bot/bot.go b/client/bot/bot.go
index 559170b..e851eb2 100644
--- a/client/bot/bot.go
+++ b/client/bot/bot.go
@@ -70,9 +70,6 @@ func Init(ctx context.Context) <-chan struct{} {
}{nil, err}
return
}
- client.API().BotsSetBotCommands(ctx, &tg.BotsSetBotCommandsRequest{
- Scope: &tg.BotCommandScopeDefault{},
- })
commands := make([]tg.BotCommand, 0, len(handlers.CommandHandlers))
for _, info := range handlers.CommandHandlers {
commands = append(commands, tg.BotCommand{Command: info.Cmd, Description: i18n.T(info.Desc)})
diff --git a/client/bot/handlers/add_task.go b/client/bot/handlers/add_task.go
index 8030a02..b089997 100644
--- a/client/bot/handlers/add_task.go
+++ b/client/bot/handlers/add_task.go
@@ -23,7 +23,11 @@ import (
)
func handleAddCallback(ctx *ext.Context, update *ext.Update) error {
- dataid := strings.Split(string(update.CallbackQuery.Data), " ")[1]
+ dataParts := strings.Split(string(update.CallbackQuery.Data), " ")
+ if len(dataParts) < 2 {
+ return fmt.Errorf("invalid callback data: %q", update.CallbackQuery.Data)
+ }
+ dataid := dataParts[1]
data, err := shortcut.GetCallbackDataWithAnswer[tcbdata.Add](ctx, update, dataid)
if err != nil {
return err
diff --git a/client/bot/handlers/cancel_task.go b/client/bot/handlers/cancel_task.go
index b886418..e3a03c1 100644
--- a/client/bot/handlers/cancel_task.go
+++ b/client/bot/handlers/cancel_task.go
@@ -1,6 +1,7 @@
package handlers
import (
+ "fmt"
"strings"
"github.com/celestix/gotgproto/dispatcher"
@@ -14,7 +15,11 @@ import (
)
func handleCancelCallback(ctx *ext.Context, update *ext.Update) error {
- taskid := strings.Split(string(update.CallbackQuery.Data), " ")[1]
+ dataParts := strings.Split(string(update.CallbackQuery.Data), " ")
+ if len(dataParts) < 2 {
+ return fmt.Errorf("invalid callback data: %q", update.CallbackQuery.Data)
+ }
+ taskid := dataParts[1]
if err := core.CancelTask(ctx, taskid); err != nil {
log.FromContext(ctx).Errorf("Failed to cancel task %s: %v", taskid, err)
ctx.AnswerCallback(msgelem.AlertCallbackAnswer(update.CallbackQuery.GetQueryID(), i18n.T(i18nk.BotMsgCancelErrorCancelFailed, map[string]any{
diff --git a/client/bot/handlers/dir.go b/client/bot/handlers/dir.go
index 5cb4f25..85058eb 100644
--- a/client/bot/handlers/dir.go
+++ b/client/bot/handlers/dir.go
@@ -10,6 +10,7 @@ import (
"github.com/krau/SaveAny-Bot/client/bot/handlers/utils/msgelem"
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
+ "github.com/krau/SaveAny-Bot/common/utils/tgutil"
"github.com/krau/SaveAny-Bot/database"
"github.com/krau/SaveAny-Bot/storage"
)
@@ -42,7 +43,9 @@ func handleDirCmd(ctx *ext.Context, update *ext.Update) error {
return dispatcher.EndGroups
}
if _, err := storage.GetStorageByUserIDAndName(ctx, user.ChatID, args[2]); err != nil {
- ctx.Reply(update, ext.ReplyTextString(err.Error()), nil)
+ logger.Errorf("Failed to get storage %q: %s", args[2], err)
+ ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgCommonErrorGetStorageFailed,
+ tgutil.EscapeHTMLTemplateData(map[string]any{"Error": err.Error()}))), nil)
return dispatcher.EndGroups
}
diff --git a/client/bot/handlers/media_group.go b/client/bot/handlers/media_group.go
index f8d870e..8230ddd 100644
--- a/client/bot/handlers/media_group.go
+++ b/client/bot/handlers/media_group.go
@@ -20,9 +20,17 @@ import (
"github.com/krau/SaveAny-Bot/storage"
)
+// mediaGroupKey uniquely identifies a media group by chat, sender, and group
+// ID so files from different users in the same chat can never be mixed.
+type mediaGroupKey struct {
+ chatID int64
+ userID int64
+ groupID int64
+}
+
type MediaGroupHandler struct {
- groups map[int64][]tfile.TGFileMessage
- timers map[int64]*time.Timer
+ groups map[mediaGroupKey][]tfile.TGFileMessage
+ timers map[mediaGroupKey]*time.Timer
mu sync.Mutex
timeout time.Duration
setupOnce sync.Once
@@ -39,8 +47,8 @@ func (m *MediaGroupHandler) SetupTimeout(timeoutSec int) {
var (
mediaGroupHandler = &MediaGroupHandler{
- groups: make(map[int64][]tfile.TGFileMessage),
- timers: make(map[int64]*time.Timer),
+ groups: make(map[mediaGroupKey][]tfile.TGFileMessage),
+ timers: make(map[mediaGroupKey]*time.Timer),
mu: sync.Mutex{},
}
)
@@ -66,32 +74,37 @@ func handleGroupMediaMessage(ctx *ext.Context, update *ext.Update, message *tg.M
}
mediaGroupHandler.mu.Lock()
defer mediaGroupHandler.mu.Unlock()
- if mediaGroupHandler.groups[groupID] == nil {
- mediaGroupHandler.groups[groupID] = make([]tfile.TGFileMessage, 0)
+ key := mediaGroupKey{
+ chatID: update.EffectiveChat().GetID(),
+ userID: userId,
+ groupID: groupID,
}
- mediaGroupHandler.groups[groupID] = append(mediaGroupHandler.groups[groupID], file)
+ if mediaGroupHandler.groups[key] == nil {
+ mediaGroupHandler.groups[key] = make([]tfile.TGFileMessage, 0)
+ }
+ mediaGroupHandler.groups[key] = append(mediaGroupHandler.groups[key], file)
- if timer, exists := mediaGroupHandler.timers[groupID]; exists {
+ if timer, exists := mediaGroupHandler.timers[key]; exists {
timer.Stop()
}
- mediaGroupHandler.timers[groupID] = time.AfterFunc(mediaGroupHandler.timeout, func() {
- processMediaGroup(ctx, update, groupID)
+ mediaGroupHandler.timers[key] = time.AfterFunc(mediaGroupHandler.timeout, func() {
+ processMediaGroup(ctx, update, key)
})
return dispatcher.EndGroups
}
-func processMediaGroup(ctx *ext.Context, update *ext.Update, groupID int64) {
+func processMediaGroup(ctx *ext.Context, update *ext.Update, key mediaGroupKey) {
logger := log.FromContext(ctx)
mediaGroupHandler.mu.Lock()
- items := mediaGroupHandler.groups[groupID]
- delete(mediaGroupHandler.groups, groupID)
- delete(mediaGroupHandler.timers, groupID)
+ items := mediaGroupHandler.groups[key]
+ delete(mediaGroupHandler.groups, key)
+ delete(mediaGroupHandler.timers, key)
mediaGroupHandler.mu.Unlock()
if len(items) == 0 {
- logger.Warn("No media items to process for group", "groupID", groupID)
+ logger.Warn("No media items to process for group", "groupID", key.groupID)
return
}
- logger.Debugf("Processing media group %d with %d items", groupID, len(items))
+ logger.Debugf("Processing media group %d with %d items", key.groupID, len(items))
userId := update.GetUserChat().GetID()
msg, err := ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgMediaGroupInfoSavingFiles, nil)), nil)
diff --git a/client/bot/handlers/middleware.go b/client/bot/handlers/middleware.go
index 1ef1dd0..3ad8079 100644
--- a/client/bot/handlers/middleware.go
+++ b/client/bot/handlers/middleware.go
@@ -22,6 +22,17 @@ func checkPermission(ctx *ext.Context, update *ext.Update) error {
return dispatcher.ContinueGroups
}
+// withPermission wraps a callback handler with the same whitelist check used
+// for message handlers (checkPermission).
+func withPermission(handler func(*ext.Context, *ext.Update) error) func(*ext.Context, *ext.Update) error {
+ return func(ctx *ext.Context, update *ext.Update) error {
+ if err := checkPermission(ctx, update); err != nil {
+ return err
+ }
+ return handler(ctx, update)
+ }
+}
+
func handleSilentMode(next func(*ext.Context, *ext.Update) error, handler func(*ext.Context, *ext.Update) error) func(*ext.Context, *ext.Update) error {
return func(ctx *ext.Context, update *ext.Update) error {
userID := update.GetUserChat().GetID()
diff --git a/client/bot/handlers/register.go b/client/bot/handlers/register.go
index 8257fa4..b7f73f5 100644
--- a/client/bot/handlers/register.go
+++ b/client/bot/handlers/register.go
@@ -56,11 +56,11 @@ func Register(disp dispatcher.Dispatcher) {
for _, info := range CommandHandlers {
disp.AddHandler(handlers.NewCommand(info.Cmd, info.handler))
}
- disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix("update"), handleUpdateCallback))
- disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeAdd), handleAddCallback))
- disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeSetDefault), handleSetDefaultCallback))
- disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeCancel), handleCancelCallback))
- disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeConfig), handleConfigCallback))
+ disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix("update"), withPermission(handleUpdateCallback)))
+ disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeAdd), withPermission(handleAddCallback)))
+ disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeSetDefault), withPermission(handleSetDefaultCallback)))
+ disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeCancel), withPermission(handleCancelCallback)))
+ disp.AddHandler(handlers.NewCallbackQuery(filters.CallbackQuery.Prefix(tcbdata.TypeConfig), withPermission(handleConfigCallback)))
disp.AddHandler(handlers.NewMessage(sabotfilters.RegexUrl(regexp.MustCompile(re.TgMessageLinkRegexString)), handleSilentMode(handleMessageLink, handleSilentSaveLink)))
disp.AddHandler(handlers.NewMessage(sabotfilters.RegexUrl(regexp.MustCompile(re.TelegraphUrlRegexString)), handleSilentMode(handleTelegraphUrlMessage, handleSilentSaveTelegraph)))
disp.AddHandler(handlers.NewMessage(filters.Message.Media, handleSilentMode(handleMediaMessage, handleSilentSaveMedia)))
diff --git a/client/bot/handlers/silent.go b/client/bot/handlers/silent.go
index fd01058..da1f8ab 100644
--- a/client/bot/handlers/silent.go
+++ b/client/bot/handlers/silent.go
@@ -1,6 +1,7 @@
package handlers
import (
+ "fmt"
"strings"
"github.com/celestix/gotgproto/dispatcher"
@@ -43,7 +44,11 @@ func handleSilentCmd(ctx *ext.Context, update *ext.Update) error {
}
func handleSetDefaultCallback(ctx *ext.Context, update *ext.Update) error {
- dataid := strings.Split(string(update.CallbackQuery.Data), " ")[1]
+ dataParts := strings.Split(string(update.CallbackQuery.Data), " ")
+ if len(dataParts) < 2 {
+ return fmt.Errorf("invalid callback data: %q", update.CallbackQuery.Data)
+ }
+ dataid := dataParts[1]
data, ok := cache.Get[tcbdata.SetDefaultStorage](dataid)
failedAnswer := func(message string) error {
diff --git a/client/bot/handlers/tasks.go b/client/bot/handlers/tasks.go
index 1d6a2d2..ac1c3c0 100644
--- a/client/bot/handlers/tasks.go
+++ b/client/bot/handlers/tasks.go
@@ -89,7 +89,11 @@ func showQueuedTasks(ctx *ext.Context, update *ext.Update) {
styling.Bold(i18n.T(i18nk.BotMsgTasksQueuedTitle)),
styling.Plain(i18n.T(i18nk.BotMsgTasksTotalPrefix, map[string]any{"Count": len(tasks)})),
)
- for _, t := range tasks {
+ const maxShown = 10
+ for i, t := range tasks {
+ if i >= maxShown {
+ break
+ }
created := t.Created.In(time.Local).Format("2006-01-02 15:04:05")
status := i18n.T(i18nk.BotMsgTasksStatusQueued)
if t.Cancelled {
@@ -105,10 +109,9 @@ func showQueuedTasks(ctx *ext.Context, update *ext.Update) {
styling.Plain("\n"+i18n.T(i18nk.BotMsgTasksFieldStatus)),
styling.Code(status),
)
- if len(tasks) > 10 {
- opts = append(opts, styling.Plain("\n"+i18n.T(i18nk.BotMsgTasksTruncatedNote, map[string]any{"Count": len(tasks)})))
- break
- }
+ }
+ if len(tasks) > maxShown {
+ opts = append(opts, styling.Plain("\n"+i18n.T(i18nk.BotMsgTasksTruncatedNote, map[string]any{"Count": len(tasks)})))
}
ctx.Reply(update, ext.ReplyTextStyledTextArray(opts), nil)
}
diff --git a/client/bot/handlers/update.go b/client/bot/handlers/update.go
index 933a9b1..52dcebc 100644
--- a/client/bot/handlers/update.go
+++ b/client/bot/handlers/update.go
@@ -10,6 +10,7 @@ import (
"github.com/celestix/gotgproto/ext"
"github.com/gotd/td/telegram/message/html"
"github.com/gotd/td/tg"
+ "github.com/krau/SaveAny-Bot/client/bot/handlers/utils/msgelem"
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
"github.com/krau/SaveAny-Bot/config"
@@ -100,7 +101,10 @@ func handleUpdateCmd(ctx *ext.Context, u *ext.Update) error {
func handleUpdateCallback(ctx *ext.Context, u *ext.Update) error {
currentV, err := semver.Parse(config.Version)
if err != nil {
- return err
+ ctx.AnswerCallback(msgelem.AlertCallbackAnswer(u.CallbackQuery.GetQueryID(), i18n.T(i18nk.BotMsgUpdateErrorVersionVarInvalid, map[string]any{
+ "Error": err.Error(),
+ })))
+ return dispatcher.EndGroups
}
ctx.EditMessage(u.GetUserChat().GetID(), &tg.MessagesEditMessageRequest{
ID: u.CallbackQuery.GetMsgID(),
diff --git a/client/user/userclient.go b/client/user/userclient.go
index cbc759f..edc9dbc 100644
--- a/client/user/userclient.go
+++ b/client/user/userclient.go
@@ -2,6 +2,7 @@ package user
import (
"context"
+ "sync"
"time"
"github.com/celestix/gotgproto"
@@ -20,17 +21,18 @@ import (
)
var uc *gotgproto.Client
-var ectx *ext.Context
+
+// getEctx lazily creates the user-client ext.Context exactly once. Guarded by
+// sync.OnceValue so concurrent GetCtx calls cannot race on ectx creation.
+var getEctx = sync.OnceValue(func() *ext.Context {
+ return uc.CreateContext()
+})
func GetCtx() *ext.Context {
- if ectx != nil {
- return ectx
- }
if uc == nil {
return nil
}
- ectx = uc.CreateContext()
- return ectx
+ return getEctx()
}
func Login(ctx context.Context) (*gotgproto.Client, error) {
diff --git a/cmd/geni18n/main.go b/cmd/geni18n/main.go
index d7cb7f0..0ba795e 100644
--- a/cmd/geni18n/main.go
+++ b/cmd/geni18n/main.go
@@ -22,7 +22,13 @@ func main() {
pkg := flag.String("pkg", "i18nk", "Package name for generated file")
flag.Parse()
+ type localeFile struct {
+ path string
+ keys map[string]struct{}
+ }
+
keys := make(map[string]struct{})
+ var localeFiles []localeFile
err := filepath.WalkDir(*dir, func(path string, d fs.DirEntry, err error) error {
if err != nil {
@@ -42,7 +48,12 @@ func main() {
return fmt.Errorf("failed to parse yaml %s: %w", path, err)
}
- collectKeys(content, "", keys)
+ fileKeys := make(map[string]struct{})
+ collectKeys(content, "", fileKeys)
+ localeFiles = append(localeFiles, localeFile{path: path, keys: fileKeys})
+ for k := range fileKeys {
+ keys[k] = struct{}{}
+ }
return nil
})
if err != nil {
@@ -50,6 +61,25 @@ func main() {
os.Exit(1)
}
+ // 一致性校验: 每个语言文件必须包含全部 key
+ invalid := false
+ for _, f := range localeFiles {
+ var missing []string
+ for k := range keys {
+ if _, ok := f.keys[k]; !ok {
+ missing = append(missing, k)
+ }
+ }
+ if len(missing) > 0 {
+ invalid = true
+ sort.Strings(missing)
+ fmt.Fprintf(os.Stderr, "Error: locale file %s is missing %d key(s): %s\n", f.path, len(missing), strings.Join(missing, ", "))
+ }
+ }
+ if invalid {
+ os.Exit(1)
+ }
+
var list []string
for k := range keys {
list = append(list, k)
diff --git a/cmd/run.go b/cmd/run.go
index b1390a2..7499296 100644
--- a/cmd/run.go
+++ b/cmd/run.go
@@ -61,6 +61,7 @@ func Run(cmd *cobra.Command, _ []string) {
<-ctx.Done()
logger.Info("Exiting...")
defer logger.Info("Exit complete")
+ core.Close()
cleanCache()
}
@@ -87,7 +88,7 @@ func initAll(ctx context.Context) (<-chan struct{}, error) {
}
}
if err := api.Start(ctx); err != nil {
- logger.Error("Failed to start API server", "error", err)
+ logger.Fatal("Failed to start API server", "error", err)
}
return bot.Init(ctx), nil
}
diff --git a/common/cache/ristretto.go b/common/cache/ristretto.go
index 5bacbf4..a54e040 100644
--- a/common/cache/ristretto.go
+++ b/common/cache/ristretto.go
@@ -2,6 +2,7 @@ package cache
import (
"fmt"
+ "sync"
"time"
"github.com/charmbracelet/log"
@@ -9,24 +10,26 @@ import (
"github.com/krau/SaveAny-Bot/config"
)
-var cache *ristretto.Cache[string, any]
+var (
+ cache *ristretto.Cache[string, any]
+ initOnce sync.Once
+)
func Init() {
- if cache != nil {
- panic("cache already initialized")
- }
- c, err := ristretto.NewCache(&ristretto.Config[string, any]{
- NumCounters: config.C().Cache.NumCounters,
- MaxCost: config.C().Cache.MaxCost,
- BufferItems: 64,
- OnReject: func(item *ristretto.Item[any]) {
- log.Warnf("Cache item rejected: key=%d, value=%v", item.Key, item.Value)
- },
+ initOnce.Do(func() {
+ c, err := ristretto.NewCache(&ristretto.Config[string, any]{
+ NumCounters: config.C().Cache.NumCounters,
+ MaxCost: config.C().Cache.MaxCost,
+ BufferItems: 64,
+ OnReject: func(item *ristretto.Item[any]) {
+ log.Warnf("Cache item rejected: key=%d, value=%v", item.Key, item.Value)
+ },
+ })
+ if err != nil {
+ log.Fatalf("failed to create ristretto cache: %v", err)
+ }
+ cache = c
})
- if err != nil {
- log.Fatalf("failed to create ristretto cache: %v", err)
- }
- cache = c
}
func Set(key string, value any) error {
diff --git a/common/i18n/i18nk/keys.go b/common/i18n/i18nk/keys.go
index 717893a..0b184b3 100644
--- a/common/i18n/i18nk/keys.go
+++ b/common/i18n/i18nk/keys.go
@@ -192,8 +192,10 @@ const (
BotMsgProgressSingleDownloadingUnknown Key = "bot.msg.progress.single_downloading_unknown"
BotMsgProgressSingleFailed Key = "bot.msg.progress.single_failed"
BotMsgProgressSingleStatusHeader Key = "bot.msg.progress.single_status_header"
- BotMsgProgressSingleUploading Key = "bot.msg.progress.single_uploading"
BotMsgProgressSingleUploadRetrying Key = "bot.msg.progress.single_upload_retrying"
+ BotMsgProgressSingleUploading Key = "bot.msg.progress.single_uploading"
+ BotMsgProgressSizeWithFiles Key = "bot.msg.progress.size_with_files"
+ BotMsgProgressSizeWithResources Key = "bot.msg.progress.size_with_resources"
BotMsgProgressTaskCanceledWithId Key = "bot.msg.progress.task_canceled_with_id"
BotMsgProgressTaskFailedWithError Key = "bot.msg.progress.task_failed_with_error"
BotMsgProgressTelegraphDonePrefix Key = "bot.msg.progress.telegraph_done_prefix"
@@ -244,7 +246,6 @@ const (
BotMsgSaveHelpText Key = "bot.msg.save_help_text"
BotMsgStorageInfoFilenamePrefix Key = "bot.msg.storage.info_filename_prefix"
BotMsgStorageInfoPromptSelectStorage Key = "bot.msg.storage.info_prompt_select_storage"
- BotMsgSyncpeersDone Key = "bot.msg.syncpeers.done"
BotMsgSyncpeersFailed Key = "bot.msg.syncpeers.failed"
BotMsgSyncpeersStart Key = "bot.msg.syncpeers.start"
BotMsgSyncpeersSuccess Key = "bot.msg.syncpeers.success"
diff --git a/common/i18n/locale/en.yaml b/common/i18n/locale/en.yaml
index 033e3f3..40f6e89 100644
--- a/common/i18n/locale/en.yaml
+++ b/common/i18n/locale/en.yaml
@@ -384,6 +384,8 @@ bot:
single_canceled: "🚫 Task canceled\n\nFilename: {{.Name}}"
single_failed: "❌ Processing failed\n\nFilename: {{.Name}}\nReason: {{.Reason}}"
downloading_prefix: "Downloading\nTotal size: "
+ size_with_files: "{{.Size}} ({{.Count}} files)"
+ size_with_resources: "{{.Size}} ({{.Count}} resources)"
processing_list_prefix: "\nProcessing:\n"
processing_none: " - None"
avg_speed_prefix: "\nAverage speed: "
@@ -424,7 +426,7 @@ bot:
transfer_failed_files_prefix: "\nFailed files: "
syncpeers:
start: "Starting to sync peers..."
- done: "Peer sync completed, total {{.Count}} chats synced"
+ success: "Peer sync completed, total {{.Count}} chats synced"
failed: "Peer sync failed: {{.Error}}"
aria2:
error_aria2_not_enabled: "Aria2 feature is not enabled in the configuration"
diff --git a/common/i18n/locale/zh-Hans.yaml b/common/i18n/locale/zh-Hans.yaml
index f688174..1b2cc71 100644
--- a/common/i18n/locale/zh-Hans.yaml
+++ b/common/i18n/locale/zh-Hans.yaml
@@ -238,9 +238,9 @@ bot:
info_install_plugin_success: "插件安装成功: {{.Name}}"
parse:
info_parsing: "正在解析..."
- error_parse_text_failed: "Failed to parse text: {{.Error}}"
- error_build_storage_select_keyboard_failed: "Failed to build storage selection keyboard: {{.Error}}"
- error_build_parsed_text_entity_failed: "Failed to build parsed text entity: {{.Error}}"
+ error_parse_text_failed: "解析文本失败: {{.Error}}"
+ error_build_storage_select_keyboard_failed: "构建存储选择键盘失败: {{.Error}}"
+ error_build_parsed_text_entity_failed: "构建解析文本实体失败: {{.Error}}"
info_link_prefix: "\n链接: "
info_author_prefix: "\n作者: "
info_description_prefix: "\n描述: "
@@ -385,6 +385,8 @@ bot:
single_canceled: "🚫 任务已取消\n\n文件名:{{.Name}}"
single_failed: "❌ 处理失败\n\n文件名:{{.Name}}\n原因:{{.Reason}}"
downloading_prefix: "正在下载\n总大小: "
+ size_with_files: "{{.Size}} ({{.Count}} 个文件)"
+ size_with_resources: "{{.Size}} ({{.Count}} 个资源)"
processing_list_prefix: "\n正在处理:\n"
processing_none: " - 无"
avg_speed_prefix: "\n平均速度: "
diff --git a/common/utils/fsutil/unique.go b/common/utils/fsutil/unique.go
new file mode 100644
index 0000000..f8b98d7
--- /dev/null
+++ b/common/utils/fsutil/unique.go
@@ -0,0 +1,27 @@
+package fsutil
+
+import (
+ "fmt"
+ "path"
+ "strings"
+
+ "github.com/rs/xid"
+)
+
+// UniquePath returns a non-taken path under basePath: the name itself, then
+// numbered variants, then a random suffix.
+func UniquePath(basePath, name string, exists func(candidate string) bool, maxAttempts int) string {
+ candidate := path.Join(basePath, name)
+ if !exists(candidate) {
+ return candidate
+ }
+ ext := path.Ext(name)
+ stem := strings.TrimSuffix(name, ext)
+ for i := 1; i <= maxAttempts; i++ {
+ candidate = path.Join(basePath, fmt.Sprintf("%s_%d%s", stem, i, ext))
+ if !exists(candidate) {
+ return candidate
+ }
+ }
+ return path.Join(basePath, fmt.Sprintf("%s_%s%s", stem, xid.New().String(), ext))
+}
diff --git a/common/utils/progressutil/progress.go b/common/utils/progressutil/progress.go
new file mode 100644
index 0000000..c33ef52
--- /dev/null
+++ b/common/utils/progressutil/progress.go
@@ -0,0 +1,50 @@
+// Package progressutil provides shared progress-update throttling for task
+// progress trackers.
+package progressutil
+
+// updatesLevels picks the percent step by file size.
+var updatesLevels = []struct {
+ size int64 // file size threshold
+ stepPercent int // minimum percent step between updates
+}{
+ {10 << 20, 100},
+ {50 << 20, 20},
+ {200 << 20, 10},
+ {500 << 20, 5},
+}
+
+// ShouldUpdate reports whether a byte-based progress update should be shown.
+func ShouldUpdate(total, downloaded int64, lastUpdatePercent int) bool {
+ if total <= 0 || downloaded <= 0 {
+ return false
+ }
+
+ percent := int((downloaded * 100) / total)
+ if percent <= lastUpdatePercent {
+ return false
+ }
+
+ step := updatesLevels[len(updatesLevels)-1].stepPercent
+ for _, lvl := range updatesLevels {
+ if total < lvl.size {
+ step = lvl.stepPercent
+ break
+ }
+ }
+
+ return percent >= lastUpdatePercent+step
+}
+
+// ShouldUpdateCount reports whether a count-based progress update (e.g. files
+// downloaded so far) should be shown: every 10 units, or when finished.
+func ShouldUpdateCount(downloaded, total int64) bool {
+ if total <= 0 || downloaded <= 0 {
+ return false
+ }
+
+ const step = int64(10)
+ if downloaded < step {
+ return downloaded == total
+ }
+ return downloaded%step == 0 || downloaded == total
+}
diff --git a/common/utils/tgutil/message.go b/common/utils/tgutil/message.go
index 0e30845..d40a6fe 100644
--- a/common/utils/tgutil/message.go
+++ b/common/utils/tgutil/message.go
@@ -194,97 +194,6 @@ func getMessagesRange(ctx *ext.Context, chatID int64, minId, maxId int) ([]*tg.M
return result, nil
}
-// [TODO]
-// type MessageItem struct {
-// Message *tg.Message
-// Error error
-// }
-
-// func IterMessages(ctx *ext.Context, chatID int64, minId, maxId int) (<-chan MessageItem, error) {
-// total := maxId - minId + 1
-// ch := make(chan MessageItem, 100)
-
-// go func() {
-// defer close(ch)
-// if !ctx.Self.Bot {
-// perr := ctx.PeerStorage.GetInputPeerById(chatID)
-// if perr == nil || perr.(*tg.InputPeerEmpty) != nil {
-// ch <- MessageItem{
-// Error: fmt.Errorf("peer not found: %d", chatID),
-// }
-// return
-// }
-
-// for i := 0; i < total; i += 100 {
-// start := minId + i
-// end := min(start+100, maxId)
-// msgs, err := ctx.Raw.MessagesGetHistory(ctx, &tg.MessagesGetHistoryRequest{
-// Peer: perr,
-// OffsetID: start,
-// AddOffset: start - end,
-// Limit: 100,
-// })
-// if err != nil {
-// ch <- MessageItem{
-// Error: fmt.Errorf("failed to get messages: %w", err),
-// }
-// return
-// }
-// var msgClass []tg.MessageClass
-// switch msgsv := msgs.(type) {
-// case *tg.MessagesMessages:
-// msgClass = msgsv.GetMessages()
-// case *tg.MessagesMessagesSlice:
-// msgClass = msgsv.GetMessages()
-// case *tg.MessagesChannelMessages:
-// msgClass = msgsv.GetMessages()
-// default:
-// ch <- MessageItem{
-// Error: fmt.Errorf("unsupported message type: %T", msgsv),
-// }
-// continue
-// }
-// for _, msg := range msgClass {
-// msg, ok := msg.AsNotEmpty()
-// if !ok {
-// continue
-// }
-// switch msg := msg.(type) {
-// case *tg.Message:
-// key := fmt.Sprintf("tgmsg:%d:%d:%d", ctx.Self.ID, chatID, msg.GetID())
-// cache.Set(key, msg)
-// ch <- MessageItem{
-// Message: msg,
-// }
-// }
-// }
-// }
-// } else {
-// for i := 0; i < total; i += 100 {
-// start := minId + i
-// end := min(start+100, maxId)
-// msgs, err := GetMessagesRange(ctx, chatID, start, end)
-// if err != nil {
-// ch <- MessageItem{
-// Error: fmt.Errorf("failed to get messages: %w", err),
-// }
-// return
-// }
-// for _, msg := range msgs {
-// if msg == nil {
-// continue
-// }
-// ch <- MessageItem{
-// Message: msg,
-// }
-// }
-// }
-// }
-// }()
-
-// return ch, nil
-// }
-
func getMessageByID(ctx *ext.Context, chatID int64, msgID int) (*tg.Message, error) {
key := fmt.Sprintf("tgmsg:%d:%d:%d", ctx.Self.ID, chatID, msgID)
if msg, ok := cache.Get[*tg.Message](key); ok {
diff --git a/config.example.toml b/config.example.toml
index ad1bd18..74dba79 100644
--- a/config.example.toml
+++ b/config.example.toml
@@ -44,6 +44,25 @@ format = ""
# 下载后转封装的视频容器格式, 留空则不转封装. 默认 mp4
recode = "mp4"
+# 解析器配置
+[parser]
+# 启用 JS 解析器插件 (Go 内置解析器默认启用)
+plugin_enable = false
+# 插件目录, 可以是多个目录
+plugin_dirs = ["./plugins"]
+# 解析器默认代理
+proxy = ""
+
+# Twitter/X 解析器配置
+[parser.twitter]
+# 自定义 API 域名
+api_domain = "api.fxtwitter.com"
+# 单独为此解析器指定代理 (留空则使用 [parser] 中的 proxy)
+# proxy = "http://127.0.0.1:7890"
+
+# Kemono 解析器配置 (暂无可配置项, 留空即可)
+[parser.kemono]
+
# HTTP API 配置
[api]
# 启用 HTTP API
diff --git a/config/hook.go b/config/hook.go
index ab01131..2c252c6 100644
--- a/config/hook.go
+++ b/config/hook.go
@@ -10,13 +10,4 @@ type hookExecConfig struct {
TaskSuccess string `toml:"task_success" mapstructure:"task_success" json:"task_success"`
TaskFail string `toml:"task_fail" mapstructure:"task_fail" json:"task_fail"`
TaskCancel string `toml:"task_cancel" mapstructure:"task_cancel" json:"task_cancel"`
-
- // TaskTypes map[string]hookExecOnTypeConfig `toml:"task_types" mapstructure:"task_types" json:"task_types"` // [TODO]
}
-
-// type hookExecOnTypeConfig struct {
-// TaskBeforeStart string `toml:"task_before_start" mapstructure:"task_before_start" json:"task_before_start"`
-// TaskSuccess string `toml:"task_success" mapstructure:"task_success" json:"task_success"`
-// TaskFail string `toml:"task_fail" mapstructure:"task_fail" json:"task_fail"`
-// TaskCancel string `toml:"task_cancel" mapstructure:"task_cancel" json:"task_cancel"`
-// }
diff --git a/config/viper.go b/config/viper.go
index 89cc831..a60d66c 100644
--- a/config/viper.go
+++ b/config/viper.go
@@ -9,6 +9,7 @@ import (
"strings"
"time"
+ "github.com/charmbracelet/log"
"github.com/duke-git/lancet/v2/slice"
"github.com/krau/SaveAny-Bot/config/storage"
"github.com/spf13/viper"
@@ -68,6 +69,13 @@ func (c Config) GetStorageByName(name string) storage.StorageConfig {
}
func Init(ctx context.Context, configFile ...string) error {
+ logger := log.FromContext(ctx)
+
+ // Reset side tables for re-init.
+ storages = nil
+ userIDs = nil
+ userStorages = make(map[int64][]string)
+
viper.SetConfigType("toml")
viper.SetEnvPrefix("SAVEANY")
viper.AutomaticEnv()
@@ -76,11 +84,13 @@ func Init(ctx context.Context, configFile ...string) error {
// 如果指定了配置文件路径,则使用指定的配置文件
// 配置文件支持传入一个 http(s) URL 地址
+ loadedFromURL := false
if len(configFile) > 0 && configFile[0] != "" {
cfg := configFile[0]
if strings.HasPrefix(cfg, "http://") || strings.HasPrefix(cfg, "https://") {
// 使用远程配置文件
- resp, err := http.Get(cfg)
+ client := &http.Client{Timeout: 30 * time.Second}
+ resp, err := client.Get(cfg)
if err != nil {
return fmt.Errorf("failed to fetch remote config file: %w", err)
}
@@ -91,6 +101,7 @@ func Init(ctx context.Context, configFile ...string) error {
if err := viper.ReadConfig(resp.Body); err != nil {
return fmt.Errorf("failed to read remote config file: %w", err)
}
+ loadedFromURL = true
} else {
viper.SetConfigFile(cfg)
}
@@ -141,13 +152,15 @@ func Init(ctx context.Context, configFile ...string) error {
viper.SetDefault(key, value)
}
- if err := viper.ReadInConfig(); err != nil {
- fmt.Println("Error reading config file, ", err)
- return err
+ if !loadedFromURL {
+ if err := viper.ReadInConfig(); err != nil {
+ logger.Errorf("Error reading config file: %v", err)
+ return err
+ }
}
if err := viper.Unmarshal(cfg); err != nil {
- fmt.Println("Error unmarshalling config file, ", err)
+ logger.Errorf("Error unmarshalling config file: %v", err)
return err
}
diff --git a/core/core.go b/core/core.go
index 0bda7c9..95561a5 100644
--- a/core/core.go
+++ b/core/core.go
@@ -3,6 +3,7 @@ package core
import (
"context"
"errors"
+ "sync"
"github.com/charmbracelet/log"
"github.com/krau/SaveAny-Bot/config"
@@ -11,7 +12,18 @@ import (
"github.com/krau/SaveAny-Bot/pkg/taskevent"
)
-var queueInstance *queue.TaskQueue[Executable]
+var (
+ queueOnce sync.Once
+ queueInstance *queue.TaskQueue[Executable]
+)
+
+// initQueue lazily creates the shared task queue.
+func initQueue() *queue.TaskQueue[Executable] {
+ queueOnce.Do(func() {
+ queueInstance = queue.NewTaskQueue[Executable]()
+ })
+ return queueInstance
+}
type Executable interface {
Type() tasktype.TaskType
@@ -65,17 +77,22 @@ func worker(ctx context.Context, qe *queue.TaskQueue[Executable], semaphore chan
func Run(ctx context.Context) {
log.FromContext(ctx).Info("Start processing tasks...")
semaphore := make(chan struct{}, config.C().Workers)
- if queueInstance == nil {
- queueInstance = queue.NewTaskQueue[Executable]()
- }
+ q := initQueue()
for range config.C().Workers {
- go worker(ctx, queueInstance, semaphore)
+ go worker(ctx, q, semaphore)
}
}
+// Close stops the queue and unblocks workers in Get.
+func Close() {
+ if q := initQueue(); q != nil {
+ q.Close()
+ }
+}
+
func AddTask(ctx context.Context, task Executable) error {
- return queueInstance.Add(queue.NewTask(ctx, task.TaskID(), task.Title(), task))
+ return initQueue().Add(queue.NewTask(ctx, task.TaskID(), task.Title(), task))
}
func CancelTask(ctx context.Context, id string) error {
diff --git a/core/tasks/batchtfile/execute.go b/core/tasks/batchtfile/execute.go
index 83ebea4..a00c602 100644
--- a/core/tasks/batchtfile/execute.go
+++ b/core/tasks/batchtfile/execute.go
@@ -2,10 +2,12 @@ package batchtfile
import (
"context"
+ "errors"
"fmt"
"io"
"os"
"path"
+ "sync"
"time"
"github.com/charmbracelet/log"
@@ -33,7 +35,9 @@ func (g executionGroup) usesBatchSaver() bool {
func (t *Task) Execute(ctx context.Context) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("batch_file[%s]", t.ID))
logger.Info("Starting batch file task")
- t.Progress.OnStart(ctx, t)
+ if t.Progress != nil {
+ t.Progress.OnStart(ctx, t)
+ }
groups := t.executionGroups()
var err error
for i := 0; i < len(groups); {
@@ -53,7 +57,11 @@ func (t *Task) Execute(ctx context.Context) error {
i = end
}
if err != nil {
- break
+ if !t.IgnoreErrors || errors.Is(err, context.Canceled) {
+ break
+ }
+ logger.Warnf("Group processing failed (ignored): %v", err)
+ err = nil
}
}
if err != nil {
@@ -62,10 +70,19 @@ func (t *Task) Execute(ctx context.Context) error {
logger.Info("Batch file task completed successfully")
}
t.finishItems(err)
- t.Progress.OnDone(ctx, t, err)
+ if t.Progress != nil {
+ t.Progress.OnDone(ctx, t, err)
+ }
return err
}
+// notifyProgress reports a progress update to the optional tracker.
+func (t *Task) notifyProgress(ctx context.Context) {
+ if t.Progress != nil {
+ t.Progress.OnProgress(ctx, t)
+ }
+}
+
func (t *Task) executionGroups() []executionGroup {
groups := make([]executionGroup, 0, len(t.elems))
for i := 0; i < len(t.elems); {
@@ -104,7 +121,13 @@ func (t *Task) processElements(ctx context.Context, elems []*TaskElement) error
return err
}
defer t.unmarkProcessing(elem.ID)
- return t.processElement(gctx, *elem)
+ err := t.processElement(gctx, *elem)
+ if err != nil && t.IgnoreErrors && !errors.Is(err, context.Canceled) {
+ // Per-item failure: keep siblings running.
+ log.FromContext(ctx).Warnf("Element %s failed (ignored): %v", elem.ID, err)
+ return nil
+ }
+ return err
})
}
return eg.Wait()
@@ -119,23 +142,51 @@ func (t *Task) processBatch(ctx context.Context, group executionGroup) error {
}
}()
+ type downloadResult struct {
+ elem *TaskElement
+ err error
+ }
+ results := make([]downloadResult, len(group.elems))
+ var resultsMu sync.Mutex
+
eg, gctx := errgroup.WithContext(ctx)
eg.SetLimit(config.C().Workers)
- for _, elem := range group.elems {
+ for i, elem := range group.elems {
eg.Go(func() error {
if err := t.markProcessing(ctx, elem); err != nil {
return err
}
defer t.unmarkProcessing(elem.ID)
- return t.downloadElement(gctx, elem)
+ err := t.downloadElement(gctx, elem)
+ // Store by original index.
+ resultsMu.Lock()
+ results[i] = downloadResult{elem: elem, err: err}
+ resultsMu.Unlock()
+ if err != nil && t.IgnoreErrors && !errors.Is(err, context.Canceled) {
+ // Per-item failure: keep siblings running.
+ log.FromContext(ctx).Warnf("Element %s failed (ignored): %v", elem.ID, err)
+ return nil
+ }
+ return err
})
}
if err := eg.Wait(); err != nil {
return err
}
- items := make([]storagetypes.BatchItem, 0, len(group.elems))
- openFiles := make([]*os.File, 0, len(group.elems))
+ // Upload only successfully downloaded elements.
+ successElems := make([]*TaskElement, 0, len(group.elems))
+ for _, r := range results {
+ if r.err == nil {
+ successElems = append(successElems, r.elem)
+ }
+ }
+ if len(successElems) == 0 {
+ return fmt.Errorf("all elements failed to download")
+ }
+
+ items := make([]storagetypes.BatchItem, 0, len(successElems))
+ openFiles := make([]*os.File, 0, len(successElems))
defer func() {
for _, file := range openFiles {
if err := file.Close(); err != nil {
@@ -143,7 +194,7 @@ func (t *Task) processBatch(ctx context.Context, group executionGroup) error {
}
}
}()
- for _, elem := range group.elems {
+ for _, elem := range successElems {
file, err := os.Open(elem.localPath)
if err != nil {
t.markItemFailed(elem.ID, FailureStageCache, err)
@@ -166,28 +217,28 @@ func (t *Task) processBatch(ctx context.Context, group executionGroup) error {
})
}
for index, item := range items {
- t.recordDownloadComplete(group.elems[index].ID, item.Size)
+ t.recordDownloadComplete(successElems[index].ID, item.Size)
}
- return t.saveBatchItems(ctx, group, items)
+ return t.saveBatchItems(ctx, successElems, items)
}
-func (t *Task) saveBatchItems(ctx context.Context, group executionGroup, items []storagetypes.BatchItem) error {
+func (t *Task) saveBatchItems(ctx context.Context, successElems []*TaskElement, items []storagetypes.BatchItem) error {
t.startUpload(ctx)
- if progressSaver, ok := group.batchSaver.(storage.StorageBatchProgressSaver); ok {
+ if progressSaver, ok := successElems[0].Storage.(storage.StorageBatchProgressSaver); ok {
err := progressSaver.SaveBatchWithProgress(ctx, items, func(index int, uploaded, total int64) {
- if index < 0 || index >= len(group.elems) {
+ if index < 0 || index >= len(successElems) {
return
}
- t.uploadCallback(ctx, group.elems[index].ID)(uploaded, total)
+ t.uploadCallback(ctx, successElems[index].ID)(uploaded, total)
})
if err != nil {
- for _, elem := range group.elems {
+ for _, elem := range successElems {
t.markItemFailed(elem.ID, FailureStageBatchUpload, err)
}
t.notifyStateChange(ctx)
return fmt.Errorf("failed to save batch: %w", err)
}
- for index, elem := range group.elems {
+ for index, elem := range successElems {
t.uploadCallback(ctx, elem.ID)(items[index].Size, items[index].Size)
t.markItemCompleted(elem.ID)
}
@@ -198,17 +249,17 @@ func (t *Task) saveBatchItems(ctx context.Context, group executionGroup, items [
items[i].Reader = ioutil.NewProgressReader(
items[i].Reader,
items[i].Size,
- t.uploadCallback(ctx, group.elems[i].ID),
+ t.uploadCallback(ctx, successElems[i].ID),
)
}
- if err := group.batchSaver.SaveBatch(ctx, items); err != nil {
- for _, elem := range group.elems {
+ if err := successElems[0].Storage.(storage.StorageBatchSaver).SaveBatch(ctx, items); err != nil {
+ for _, elem := range successElems {
t.markItemFailed(elem.ID, FailureStageBatchUpload, err)
}
t.notifyStateChange(ctx)
return fmt.Errorf("failed to save batch: %w", err)
}
- for index, elem := range group.elems {
+ for index, elem := range successElems {
t.uploadCallback(ctx, elem.ID)(items[index].Size, items[index].Size)
t.markItemCompleted(elem.ID)
}
@@ -225,7 +276,7 @@ func (t *Task) markProcessing(ctx context.Context, elem *TaskElement) error {
t.processing[elem.ID] = elem
t.processingMu.Unlock()
t.markItemActive(elem.ID, elem.stream, time.Now())
- t.Progress.OnProgress(ctx, t)
+ t.notifyProgress(ctx)
return nil
}
@@ -247,7 +298,7 @@ func (t *Task) downloadElement(ctx context.Context, elem *TaskElement) error {
wrAt := ioutil.NewProgressWriterAt(localFile, func(n int) {
t.recordItemDownload(elem.ID, int64(n), time.Now())
downloaded := t.downloaded.Add(int64(n))
- t.Progress.OnProgress(ctx, t)
+ t.notifyProgress(ctx)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
@@ -274,7 +325,7 @@ func (t *Task) downloadElement(ctx context.Context, elem *TaskElement) error {
}
}
t.markItemDownloaded(elem.ID)
- t.Progress.OnProgress(ctx, t)
+ t.notifyProgress(ctx)
return nil
}
@@ -295,7 +346,7 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
wr := ioutil.NewProgressWriter(pw, func(n int) {
t.recordItemDownload(elem.ID, int64(n), time.Now())
downloaded := t.downloaded.Add(int64(n))
- t.Progress.OnProgress(ctx, t)
+ t.notifyProgress(ctx)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
@@ -318,7 +369,12 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
if err := errg.Wait(); err != nil {
return fmt.Errorf("failed to download file in stream mode: %w", err)
}
- t.recordDownloadComplete(elem.ID, 0)
+ // Streamed bytes are the uploaded bytes.
+ var streamedBytes int64
+ t.updateItem(elem.ID, func(item *itemProgressState) {
+ streamedBytes = item.downloaded
+ })
+ t.recordDownloadComplete(elem.ID, streamedBytes)
t.markItemCompleted(elem.ID)
t.notifyStateChange(ctx)
logger.Info("File downloaded successfully in stream mode")
@@ -339,7 +395,7 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
wrAt := ioutil.NewProgressWriterAt(localFile, func(n int) {
t.recordItemDownload(elem.ID, int64(n), time.Now())
downloaded := t.downloaded.Add(int64(n))
- t.Progress.OnProgress(ctx, t)
+ t.notifyProgress(ctx)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
diff --git a/core/tasks/batchtfile/notify_test.go b/core/tasks/batchtfile/notify_test.go
new file mode 100644
index 0000000..0a27c4b
--- /dev/null
+++ b/core/tasks/batchtfile/notify_test.go
@@ -0,0 +1,35 @@
+package batchtfile
+
+import (
+ "context"
+ "sync/atomic"
+ "testing"
+)
+
+type recordingTracker struct {
+ calls atomic.Int64
+}
+
+func (r *recordingTracker) OnStart(context.Context, TaskInfo) {}
+func (r *recordingTracker) OnDone(context.Context, TaskInfo, error) {
+}
+func (r *recordingTracker) OnProgress(context.Context, TaskInfo) {
+ r.calls.Add(1)
+}
+
+// Regression: notifyProgress must call the tracker, not itself. The previous
+// self-call recursed until stack overflow on any batch task with a tracker.
+func TestNotifyProgressCallsTracker(t *testing.T) {
+ tracker := &recordingTracker{}
+ task := &Task{Progress: tracker}
+ task.notifyProgress(t.Context())
+ if tracker.calls.Load() != 1 {
+ t.Fatalf("expected 1 tracker call, got %d", tracker.calls.Load())
+ }
+}
+
+// The nil tracker path must stay a no-op.
+func TestNotifyProgressNilTracker(t *testing.T) {
+ task := &Task{}
+ task.notifyProgress(t.Context()) // must not panic
+}
diff --git a/core/tasks/batchtfile/progress.go b/core/tasks/batchtfile/progress.go
index eb96dc1..52fac1d 100644
--- a/core/tasks/batchtfile/progress.go
+++ b/core/tasks/batchtfile/progress.go
@@ -194,10 +194,13 @@ func buildBatchDoneMarkup(info TaskInfo, skipped []string, err error) string {
totalSize = info.TotalSize()
}
if err == nil {
- if len(skipped) > 0 {
+ completed, _, _, failed := itemCounts(items)
+ // Report per-element failures instead of full completion.
+ totalSkipped := len(skipped) + failed
+ if totalSkipped > 0 {
return localizedProgressMarkup(i18nk.BotMsgProgressBatchDoneWithSkipped, map[string]any{
- "Success": len(items),
- "Skipped": len(skipped),
+ "Success": completed,
+ "Skipped": totalSkipped,
"Size": dlutil.FormatSize(totalSize),
})
}
diff --git a/core/tasks/batchtfile/task.go b/core/tasks/batchtfile/task.go
index 5d0472a..ca05f52 100644
--- a/core/tasks/batchtfile/task.go
+++ b/core/tasks/batchtfile/task.go
@@ -47,7 +47,6 @@ type Task struct {
uploadOnce sync.Once
uploadMu sync.Mutex
uploaded map[string]int64
- failed map[string]error // [TODO] errors for each element
}
// Title implements core.Exectable.
@@ -136,7 +135,6 @@ func NewBatchTGFileTask(
uploaded: make(map[string]int64),
IgnoreErrors: ignoreErrors,
processingMu: sync.RWMutex{},
- failed: make(map[string]error),
}
return task
}
diff --git a/core/tasks/batchtfile/utils.go b/core/tasks/batchtfile/utils.go
deleted file mode 100644
index b8c39cd..0000000
--- a/core/tasks/batchtfile/utils.go
+++ /dev/null
@@ -1,32 +0,0 @@
-package batchtfile
-
-var progressUpdatesLevels = []struct {
- size int64 // 文件大小阈值
- stepPercent int // 每多少 % 更新一次
-}{
- {10 << 20, 100},
- {50 << 20, 20},
- {200 << 20, 10},
- {500 << 20, 5},
-}
-
-func shouldUpdateProgress(total, downloaded int64, lastUpdatePercent int) bool {
- if total <= 0 || downloaded <= 0 {
- return false
- }
-
- percent := int((downloaded * 100) / total)
- if percent <= lastUpdatePercent {
- return false
- }
-
- step := progressUpdatesLevels[len(progressUpdatesLevels)-1].stepPercent
- for _, lvl := range progressUpdatesLevels {
- if total < lvl.size {
- step = lvl.stepPercent
- break
- }
- }
-
- return percent >= lastUpdatePercent+step
-}
diff --git a/core/tasks/directlinks/execute.go b/core/tasks/directlinks/execute.go
index 4ea50cc..d4a6cf3 100644
--- a/core/tasks/directlinks/execute.go
+++ b/core/tasks/directlinks/execute.go
@@ -76,12 +76,11 @@ func (t *Task) Execute(ctx context.Context) error {
eg.SetLimit(config.C().Workers)
for _, file := range t.files {
eg.Go(func() error {
- t.processingMu.RLock()
+ t.processingMu.Lock()
if _, ok := t.processing[file.URL]; ok {
+ t.processingMu.Unlock()
return fmt.Errorf("file %s is already being processed", file.URL)
}
- t.processingMu.RUnlock()
- t.processingMu.Lock()
t.processing[file.URL] = file
t.processingMu.Unlock()
defer func() {
@@ -90,7 +89,6 @@ func (t *Task) Execute(ctx context.Context) error {
t.processingMu.Unlock()
}()
err := t.processLink(gctx, file)
- t.downloaded.Add(1)
if errors.Is(err, context.Canceled) {
logger.Debug("Link processing canceled")
return err
@@ -99,6 +97,7 @@ func (t *Task) Execute(ctx context.Context) error {
logger.Errorf("Error processing link %s: %v", file.URL, err)
return fmt.Errorf("failed to process link %s: %w", file.URL, err)
}
+ t.downloaded.Add(1)
return nil
})
}
diff --git a/core/tasks/directlinks/progress.go b/core/tasks/directlinks/progress.go
index 223a50c..bd72ac2 100644
--- a/core/tasks/directlinks/progress.go
+++ b/core/tasks/directlinks/progress.go
@@ -15,6 +15,7 @@ import (
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
+ "github.com/krau/SaveAny-Bot/common/utils/progressutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
)
@@ -102,7 +103,7 @@ func (p *Progress) OnDone(ctx context.Context, info TaskInfo, err error) {
// OnProgress implements ProgressTracker.
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
- if !shouldUpdateProgress(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
+ if !progressutil.ShouldUpdate(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
return
}
percent := int((info.DownloadedBytes() * 100) / info.TotalBytes())
@@ -115,7 +116,10 @@ func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
var entities []tg.MessageEntityClass
if err := styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressDownloadingPrefix, nil)),
- styling.Code(fmt.Sprintf("%.2f MB (%d个文件)", float64(info.TotalBytes())/(1024*1024), info.TotalFiles())),
+ styling.Code(i18n.T(i18nk.BotMsgProgressSizeWithFiles, map[string]any{
+ "Size": fmt.Sprintf("%.2f MB", float64(info.TotalBytes())/(1024*1024)),
+ "Count": info.TotalFiles(),
+ })),
styling.Plain(i18n.T(i18nk.BotMsgProgressProcessingListPrefix, nil)),
func() styling.StyledTextOption {
var lines []string
diff --git a/core/tasks/directlinks/task.go b/core/tasks/directlinks/task.go
index 0ffc793..9e4e08f 100644
--- a/core/tasks/directlinks/task.go
+++ b/core/tasks/directlinks/task.go
@@ -45,7 +45,6 @@ type Task struct {
downloaded atomic.Int64 // downloaded files count
processing map[string]*File // {"url": File}
processingMu sync.RWMutex
- failed map[string]error // [TODO] errors for each file
}
// Title implements core.Exectable.
@@ -127,7 +126,6 @@ func NewTask(
client: http.DefaultClient,
processing: make(map[string]*File),
processingMu: sync.RWMutex{},
- failed: make(map[string]error),
totalFiles: int64(len(files)),
}
}
diff --git a/core/tasks/directlinks/util.go b/core/tasks/directlinks/util.go
index 57d19c2..5373b9f 100644
--- a/core/tasks/directlinks/util.go
+++ b/core/tasks/directlinks/util.go
@@ -207,34 +207,3 @@ func parseFilenameFallback(cd string) string {
return decodeFilenameParam(value)
}
-
-var progressUpdatesLevels = []struct {
- size int64 // 文件大小阈值
- stepPercent int // 每多少 % 更新一次
-}{
- {10 << 20, 100},
- {50 << 20, 50},
- {200 << 20, 20},
- {500 << 20, 10},
-}
-
-func shouldUpdateProgress(total, downloaded int64, lastUpdatePercent int) bool {
- if total <= 0 || downloaded <= 0 {
- return false
- }
-
- percent := int((downloaded * 100) / total)
- if percent <= lastUpdatePercent {
- return false
- }
-
- step := progressUpdatesLevels[len(progressUpdatesLevels)-1].stepPercent
- for _, lvl := range progressUpdatesLevels {
- if total < lvl.size {
- step = lvl.stepPercent
- break
- }
- }
-
- return percent >= lastUpdatePercent+step
-}
diff --git a/core/tasks/parsed/execute.go b/core/tasks/parsed/execute.go
index 9075e81..18bf3b5 100644
--- a/core/tasks/parsed/execute.go
+++ b/core/tasks/parsed/execute.go
@@ -30,21 +30,20 @@ func (t *Task) Execute(ctx context.Context) error {
eg.SetLimit(config.C().Workers)
for _, resource := range t.item.Resources {
eg.Go(func() error {
- t.processingMu.RLock()
- if t.processing[resource.ID()] != nil {
- return fmt.Errorf("resource %s is already being processed", resource.ID())
- }
- t.processingMu.RUnlock()
+ resourceID := resource.ID()
t.processingMu.Lock()
- t.processing[resource.ID()] = &resource
+ if t.processing[resourceID] != nil {
+ t.processingMu.Unlock()
+ return fmt.Errorf("resource %s is already being processed", resourceID)
+ }
+ t.processing[resourceID] = &resource
t.processingMu.Unlock()
defer func() {
t.processingMu.Lock()
- delete(t.processing, resource.URL)
+ delete(t.processing, resourceID)
t.processingMu.Unlock()
}()
err := t.processResource(gctx, resource)
- t.downloaded.Add(1)
if errors.Is(err, context.Canceled) {
logger.Debug("Resource processing canceled")
return err
@@ -53,6 +52,7 @@ func (t *Task) Execute(ctx context.Context) error {
logger.Errorf("Error processing resource %s: %v", resource.URL, err)
return fmt.Errorf("failed to process resource %s: %w", resource.URL, err)
}
+ t.downloaded.Add(1)
return nil
})
}
diff --git a/core/tasks/parsed/progress.go b/core/tasks/parsed/progress.go
index 0cac6c4..66debb3 100644
--- a/core/tasks/parsed/progress.go
+++ b/core/tasks/parsed/progress.go
@@ -15,40 +15,10 @@ import (
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
+ "github.com/krau/SaveAny-Bot/common/utils/progressutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
)
-var progressUpdatesLevels = []struct {
- size int64 // 文件大小阈值
- stepPercent int // 每多少 % 更新一次
-}{
- {10 << 20, 100},
- {50 << 20, 50},
- {200 << 20, 20},
- {500 << 20, 10},
-}
-
-func shouldUpdateProgress(total, downloaded int64, lastUpdatePercent int) bool {
- if total <= 0 || downloaded <= 0 {
- return false
- }
-
- percent := int((downloaded * 100) / total)
- if percent <= lastUpdatePercent {
- return false
- }
-
- step := progressUpdatesLevels[len(progressUpdatesLevels)-1].stepPercent
- for _, lvl := range progressUpdatesLevels {
- if total < lvl.size {
- step = lvl.stepPercent
- break
- }
- }
-
- return percent >= lastUpdatePercent+step
-}
-
type ProgressTracker interface {
OnStart(ctx context.Context, info TaskInfo)
OnProgress(ctx context.Context, info TaskInfo)
@@ -73,7 +43,10 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
styling.Plain(i18n.T(i18nk.BotMsgProgressParsedStartPrefix, map[string]any{
"Site": info.Site(),
})),
- styling.Code(fmt.Sprintf("%.2f MB (%d个资源)", float64(info.TotalBytes())/(1024*1024), info.TotalResources())),
+ styling.Code(i18n.T(i18nk.BotMsgProgressSizeWithResources, map[string]any{
+ "Size": fmt.Sprintf("%.2f MB", float64(info.TotalBytes())/(1024*1024)),
+ "Count": info.TotalResources(),
+ })),
); err != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", err)
return
@@ -101,7 +74,7 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
- if !shouldUpdateProgress(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
+ if !progressutil.ShouldUpdate(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
return
}
percent := int((info.DownloadedBytes() * 100) / info.TotalBytes())
@@ -114,7 +87,10 @@ func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
var entities []tg.MessageEntityClass
if err := styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressDownloadingPrefix, nil)),
- styling.Code(fmt.Sprintf("%.2f MB (%d个文件)", float64(info.TotalBytes())/(1024*1024), info.TotalResources())),
+ styling.Code(i18n.T(i18nk.BotMsgProgressSizeWithResources, map[string]any{
+ "Size": fmt.Sprintf("%.2f MB", float64(info.TotalBytes())/(1024*1024)),
+ "Count": info.TotalResources(),
+ })),
styling.Plain(i18n.T(i18nk.BotMsgProgressProcessingListPrefix, nil)),
func() styling.StyledTextOption {
var lines []string
diff --git a/core/tasks/parsed/task.go b/core/tasks/parsed/task.go
index 737830f..32b7b49 100644
--- a/core/tasks/parsed/task.go
+++ b/core/tasks/parsed/task.go
@@ -33,7 +33,6 @@ type Task struct {
downloadedBytes atomic.Int64 // downloaded bytes count
processing map[string]ResourceInfo
processingMu sync.RWMutex
- failed map[string]error // [TODO] errors for each resource
}
// Title implements core.Exectable.
@@ -84,6 +83,5 @@ func NewTask(
progress: progressTracker,
processing: make(map[string]ResourceInfo),
processingMu: sync.RWMutex{},
- failed: make(map[string]error),
}
}
diff --git a/core/tasks/telegraph/execute.go b/core/tasks/telegraph/execute.go
index 0d67803..d6d891a 100644
--- a/core/tasks/telegraph/execute.go
+++ b/core/tasks/telegraph/execute.go
@@ -18,7 +18,9 @@ import (
func (t *Task) Execute(ctx context.Context) error {
logger := log.FromContext(ctx)
logger.Infof("Starting Telegraph task %s", t.PhPath)
- t.progress.OnStart(ctx, t)
+ if t.progress != nil {
+ t.progress.OnStart(ctx, t)
+ }
eg, gctx := errgroup.WithContext(ctx)
eg.SetLimit(config.C().Workers)
for i, pic := range t.Pics {
@@ -29,7 +31,9 @@ func (t *Task) Execute(ctx context.Context) error {
return fmt.Errorf("failed to process picture %s: %w", pic, err)
}
downloaded := t.downloaded.Add(1)
- t.progress.OnProgress(gctx, t)
+ if t.progress != nil {
+ t.progress.OnProgress(gctx, t)
+ }
taskevent.Emit(gctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
@@ -45,7 +49,9 @@ func (t *Task) Execute(ctx context.Context) error {
} else {
logger.Infof("Telegraph task %s completed successfully", t.PhPath)
}
- t.progress.OnDone(ctx, t, err)
+ if t.progress != nil {
+ t.progress.OnDone(ctx, t, err)
+ }
return err
}
diff --git a/core/tasks/telegraph/nil_tracker_test.go b/core/tasks/telegraph/nil_tracker_test.go
new file mode 100644
index 0000000..984ff49
--- /dev/null
+++ b/core/tasks/telegraph/nil_tracker_test.go
@@ -0,0 +1,31 @@
+package telegraph_test
+
+import (
+ "context"
+ "io"
+ "testing"
+
+ storconfig "github.com/krau/SaveAny-Bot/config/storage"
+ "github.com/krau/SaveAny-Bot/core/tasks/telegraph"
+ storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
+ "github.com/krau/SaveAny-Bot/storage"
+)
+
+type mockStorage struct{}
+
+func (mockStorage) Init(context.Context, storconfig.StorageConfig) error { return nil }
+func (mockStorage) Type() storenum.StorageType { return storenum.Local }
+func (mockStorage) Name() string { return "mock" }
+func (mockStorage) Save(context.Context, io.Reader, string) error { return nil }
+func (mockStorage) Exists(context.Context, string) bool { return false }
+
+var _ storage.Storage = mockStorage{}
+
+// Regression: API-created tasks run with a nil ProgressTracker (progress goes
+// through taskevent); Execute must not panic on the tracker callbacks.
+func TestExecuteWithNilTracker(t *testing.T) {
+ task := telegraph.NewTask("id", t.Context(), "/page", nil, mockStorage{}, "/out", nil, nil)
+ if err := task.Execute(t.Context()); err != nil {
+ t.Fatalf("Execute returned error: %v", err)
+ }
+}
diff --git a/core/tasks/telegraph/progress.go b/core/tasks/telegraph/progress.go
index cc58027..bac2d50 100644
--- a/core/tasks/telegraph/progress.go
+++ b/core/tasks/telegraph/progress.go
@@ -11,6 +11,7 @@ import (
"github.com/gotd/td/tg"
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
+ "github.com/krau/SaveAny-Bot/common/utils/progressutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
)
@@ -60,7 +61,7 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
- if !shouldUpdateProgress(info.Downloaded(), int64(info.TotalPics())) {
+ if !progressutil.ShouldUpdateCount(info.Downloaded(), int64(info.TotalPics())) {
return
}
log.FromContext(ctx).Debugf("Progress update: %s, %d/%d", info.TaskID(), info.Downloaded(), info.TotalPics())
diff --git a/core/tasks/telegraph/utils.go b/core/tasks/telegraph/utils.go
deleted file mode 100644
index ce109eb..0000000
--- a/core/tasks/telegraph/utils.go
+++ /dev/null
@@ -1,13 +0,0 @@
-package telegraph
-
-func shouldUpdateProgress(downloaded int64, total int64) bool {
- if total <= 0 || downloaded <= 0 {
- return false
- }
-
- step := int64(10)
- if downloaded < step {
- return downloaded == total
- }
- return downloaded%step == 0 || downloaded == total
-}
diff --git a/core/tasks/tfile/progress.go b/core/tasks/tfile/progress.go
index fb6be21..c0bbe91 100644
--- a/core/tasks/tfile/progress.go
+++ b/core/tasks/tfile/progress.go
@@ -14,6 +14,7 @@ import (
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
+ "github.com/krau/SaveAny-Bot/common/utils/progressutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
)
@@ -111,7 +112,7 @@ func (p *Progress) OnProgress(ctx context.Context, info TaskInfo, downloaded, to
func shouldUpdateSingleDownloadProgress(total, downloaded int64, lastPercent int, elapsed time.Duration) bool {
if total > 0 {
- return shouldUpdateProgress(total, downloaded, lastPercent)
+ return progressutil.ShouldUpdate(total, downloaded, lastPercent)
}
return downloaded > 0 && elapsed >= uploadProgressMaxInterval
}
@@ -183,7 +184,7 @@ func shouldUpdateUploadProgress(total, uploaded int64, lastPercent int, elapsed
if percent == lastPercent {
return elapsed >= uploadProgressMaxInterval
}
- return shouldUpdateProgress(total, uploaded, lastPercent) || elapsed >= uploadProgressMaxInterval
+ return progressutil.ShouldUpdate(total, uploaded, lastPercent) || elapsed >= uploadProgressMaxInterval
}
func singleUploadPhase(attempt int) singleProgressPhase {
diff --git a/core/tasks/tfile/util.go b/core/tasks/tfile/util.go
deleted file mode 100644
index 81862d9..0000000
--- a/core/tasks/tfile/util.go
+++ /dev/null
@@ -1,32 +0,0 @@
-package tfile
-
-var progressUpdatesLevels = []struct {
- size int64 // 文件大小阈值
- stepPercent int // 每多少 % 更新一次
-}{
- {10 << 20, 100},
- {50 << 20, 20},
- {200 << 20, 10},
- {500 << 20, 5},
-}
-
-func shouldUpdateProgress(total, downloaded int64, lastUpdatePercent int) bool {
- if total <= 0 || downloaded <= 0 {
- return false
- }
-
- percent := int((downloaded * 100) / total)
- if percent <= lastUpdatePercent {
- return false
- }
-
- step := progressUpdatesLevels[len(progressUpdatesLevels)-1].stepPercent
- for _, lvl := range progressUpdatesLevels {
- if total < lvl.size {
- step = lvl.stepPercent
- break
- }
- }
-
- return percent >= lastUpdatePercent+step
-}
diff --git a/core/tasks/transfer/cancel_test.go b/core/tasks/transfer/cancel_test.go
new file mode 100644
index 0000000..7675d8e
--- /dev/null
+++ b/core/tasks/transfer/cancel_test.go
@@ -0,0 +1,75 @@
+package transfer_test
+
+import (
+ "context"
+ "errors"
+ "io"
+ "os"
+ "path/filepath"
+ "testing"
+
+ "github.com/krau/SaveAny-Bot/config"
+ storconfig "github.com/krau/SaveAny-Bot/config/storage"
+ "github.com/krau/SaveAny-Bot/core/tasks/transfer"
+ storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
+ "github.com/krau/SaveAny-Bot/pkg/storagetypes"
+ "github.com/krau/SaveAny-Bot/storage"
+)
+
+// initConfig seeds the global config so task execution reads a sane Workers
+// value (the zero default would deadlock errgroup.SetLimit(0)).
+func initConfig(t *testing.T) {
+ t.Helper()
+ path := filepath.Join(t.TempDir(), "config.toml")
+ if err := os.WriteFile(path, []byte("workers = 2\n"), 0o644); err != nil {
+ t.Fatal(err)
+ }
+ if err := config.Init(t.Context(), path); err != nil {
+ t.Fatal(err)
+ }
+}
+
+type cancelSource struct{}
+
+func (cancelSource) Init(context.Context, storconfig.StorageConfig) error { return nil }
+func (cancelSource) Type() storenum.StorageType { return storenum.Local }
+func (cancelSource) Name() string { return "cancel-source" }
+func (cancelSource) Save(context.Context, io.Reader, string) error { return nil }
+func (cancelSource) Exists(context.Context, string) bool { return false }
+func (cancelSource) ListFiles(context.Context, string) ([]storagetypes.FileInfo, error) {
+ return nil, nil
+}
+
+// OpenFile reports cancellation, as the task context would be cancelled.
+func (cancelSource) OpenFile(ctx context.Context, path string) (io.ReadCloser, int64, error) {
+ return nil, 0, ctx.Err()
+}
+
+type voidTarget struct{}
+
+func (voidTarget) Init(context.Context, storconfig.StorageConfig) error { return nil }
+func (voidTarget) Type() storenum.StorageType { return storenum.Local }
+func (voidTarget) Name() string { return "void-target" }
+func (voidTarget) Save(context.Context, io.Reader, string) error { return nil }
+func (voidTarget) Exists(context.Context, string) bool { return false }
+
+var (
+ _ storage.StorageReadable = cancelSource{}
+ _ storage.Storage = voidTarget{}
+)
+
+// Regression: IgnoreErrors must swallow element failures but never a task
+// cancellation; Execute must surface context.Canceled.
+func TestIgnoreErrorsPropagatesCancel(t *testing.T) {
+ initConfig(t)
+ ctx, cancel := context.WithCancel(t.Context())
+ cancel()
+
+ elem := transfer.NewTaskElement(cancelSource{}, storagetypes.FileInfo{Path: "/a", Name: "a", Size: 1}, voidTarget{}, "/out")
+ task := transfer.NewTransferTask("id", ctx, []transfer.TaskElement{*elem}, nil, true)
+
+ err := task.Execute(ctx)
+ if !errors.Is(err, context.Canceled) {
+ t.Fatalf("expected context.Canceled, got %v", err)
+ }
+}
diff --git a/core/tasks/transfer/execute.go b/core/tasks/transfer/execute.go
index 13cc524..b517f05 100644
--- a/core/tasks/transfer/execute.go
+++ b/core/tasks/transfer/execute.go
@@ -2,6 +2,7 @@ package transfer
import (
"context"
+ "errors"
"fmt"
"io"
"os"
@@ -20,7 +21,9 @@ import (
func (t *Task) Execute(ctx context.Context) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("transfer[%s]", t.ID))
logger.Info("Starting transfer task")
- t.Progress.OnStart(ctx, t)
+ if t.Progress != nil {
+ t.Progress.OnStart(ctx, t)
+ }
workers := config.C().Workers
eg, gctx := errgroup.WithContext(ctx)
@@ -28,14 +31,11 @@ func (t *Task) Execute(ctx context.Context) error {
for _, elem := range t.elems {
eg.Go(func() error {
- t.processingMu.RLock()
+ t.processingMu.Lock()
if t.processing[elem.ID] != nil {
- t.processingMu.RUnlock()
+ t.processingMu.Unlock()
return fmt.Errorf("element with ID %s is already being processed", elem.ID)
}
- t.processingMu.RUnlock()
-
- t.processingMu.Lock()
t.processing[elem.ID] = &elem
t.processingMu.Unlock()
@@ -46,7 +46,7 @@ func (t *Task) Execute(ctx context.Context) error {
}()
err := t.processElement(gctx, elem)
- if err != nil && !t.IgnoreErrors {
+ if err != nil && (!t.IgnoreErrors || errors.Is(err, context.Canceled)) {
return err
}
if err != nil {
@@ -66,7 +66,9 @@ func (t *Task) Execute(ctx context.Context) error {
logger.Info("Transfer task completed successfully")
}
- t.Progress.OnDone(ctx, t, err)
+ if t.Progress != nil {
+ t.Progress.OnDone(ctx, t, err)
+ }
return err
}
@@ -116,7 +118,9 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
}
t.uploaded.Add(size)
- t.Progress.OnProgress(ctx, t)
+ if t.Progress != nil {
+ t.Progress.OnProgress(ctx, t)
+ }
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
diff --git a/core/tasks/transfer/progress.go b/core/tasks/transfer/progress.go
index 2852524..bd538b3 100644
--- a/core/tasks/transfer/progress.go
+++ b/core/tasks/transfer/progress.go
@@ -14,6 +14,7 @@ import (
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
+ "github.com/krau/SaveAny-Bot/common/utils/progressutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
)
@@ -83,7 +84,7 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
- if !shouldUpdateProgress(info.TotalSize(), info.Uploaded(), int(p.lastUpdatePercent.Load())) {
+ if !progressutil.ShouldUpdate(info.TotalSize(), info.Uploaded(), int(p.lastUpdatePercent.Load())) {
return
}
percent := int((info.Uploaded() * 100) / info.TotalSize())
@@ -221,14 +222,6 @@ func (p *Progress) OnDone(ctx context.Context, info TaskInfo, err error) {
}
}
-func shouldUpdateProgress(total, current int64, lastPercent int) bool {
- if total == 0 {
- return false
- }
- currentPercent := int((current * 100) / total)
- return currentPercent > lastPercent && currentPercent%5 == 0
-}
-
func formatDuration(d time.Duration) string {
d = d.Round(time.Second)
h := d / time.Hour
diff --git a/parsers/js/api.go b/parsers/js/api.go
index 6cd28ae..1ae6597 100644
--- a/parsers/js/api.go
+++ b/parsers/js/api.go
@@ -44,7 +44,10 @@ func jsRegisterParser(vm *goja.Runtime) func(call goja.FunctionCall) goja.Value
return vm.NewGoError(errors.New("metadata cannot be null or undefined"))
}
- pluginV := semver.MustParse(metadata.Version)
+ pluginV, err := semver.Parse(metadata.Version)
+ if err != nil {
+ return vm.NewGoError(fmt.Errorf("invalid parser version %q: %w", metadata.Version, err))
+ }
if pluginV.LT(MinimumParserVersion) {
return vm.NewGoError(fmt.Errorf("parser version %s is not supported, must be at least %s", metadata.Version, MinimumParserVersion))
}
@@ -57,6 +60,9 @@ func jsRegisterParser(vm *goja.Runtime) func(call goja.FunctionCall) goja.Value
if parseFn == nil || goja.IsUndefined(parseFn) {
return vm.NewGoError(errors.New("parser must provide a parse function"))
}
+ if handleFn == nil || goja.IsUndefined(handleFn) {
+ return vm.NewGoError(errors.New("parser must provide a canHandle function"))
+ }
parsers.Add(newJSParser(vm, handleFn, parseFn, metadata))
return goja.Undefined()
}
diff --git a/parsers/js/js.go b/parsers/js/js.go
index 32629cb..4b21ad4 100644
--- a/parsers/js/js.go
+++ b/parsers/js/js.go
@@ -9,6 +9,7 @@ import (
"os"
"path/filepath"
"sync"
+ "time"
"github.com/charmbracelet/log"
"github.com/dop251/goja"
@@ -34,11 +35,23 @@ type jsParserResp struct {
err error
}
+const canHandleTimeout = 10 * time.Second
+
func (p *jsParser) CanHandle(url string) bool {
respCh := make(chan jsParserResp, 1)
- p.reqCh <- jsParserReq{method: ParserMethodCanHandle, url: url, respCh: respCh}
- resp := <-respCh
- return resp.ok && resp.err == nil
+ timer := time.NewTimer(canHandleTimeout)
+ defer timer.Stop()
+ select {
+ case p.reqCh <- jsParserReq{method: ParserMethodCanHandle, url: url, respCh: respCh}:
+ case <-timer.C:
+ return false
+ }
+ select {
+ case resp := <-respCh:
+ return resp.ok && resp.err == nil
+ case <-timer.C:
+ return false
+ }
}
func (p *jsParser) Parse(ctx context.Context, url string) (*parser.Item, error) {
@@ -61,41 +74,57 @@ func newJSParser(vm *goja.Runtime, canHandleFunc, parseFunc goja.Value, metadata
go func() {
for req := range p.reqCh {
- switch req.method {
- case ParserMethodCanHandle:
- fn, _ := goja.AssertFunction(canHandleFunc)
- res, err := fn(goja.Undefined(), p.vm.ToValue(req.url))
- if err != nil {
- req.respCh <- jsParserResp{ok: false, err: err}
- continue
- }
- req.respCh <- jsParserResp{ok: res.ToBoolean()}
- case ParserMethodParse:
- fn, _ := goja.AssertFunction(parseFunc)
- result, err := fn(goja.Undefined(), p.vm.ToValue(req.url))
- if err != nil {
- req.respCh <- jsParserResp{err: err}
- continue
- }
-
- var item parser.Item
- if exported := result.Export(); exported != nil {
- data, err := json.Marshal(exported)
+ func() {
+ defer func() {
+ if r := recover(); r != nil {
+ log.Errorf("JS parser %q panicked while handling method %d: %v", p.meta.Name, req.method, r)
+ req.respCh <- jsParserResp{err: fmt.Errorf("JS parser %q panicked: %v", p.meta.Name, r)}
+ }
+ }()
+ switch req.method {
+ case ParserMethodCanHandle:
+ fn, ok := goja.AssertFunction(canHandleFunc)
+ if !ok {
+ req.respCh <- jsParserResp{err: fmt.Errorf("canHandle is not a function")}
+ return
+ }
+ res, err := fn(goja.Undefined(), p.vm.ToValue(req.url))
if err != nil {
- req.respCh <- jsParserResp{err: fmt.Errorf("failed to marshal result to JSON: %w", err)}
- continue
+ req.respCh <- jsParserResp{ok: false, err: err}
+ return
+ }
+ req.respCh <- jsParserResp{ok: res.ToBoolean()}
+ case ParserMethodParse:
+ fn, ok := goja.AssertFunction(parseFunc)
+ if !ok {
+ req.respCh <- jsParserResp{err: fmt.Errorf("parse is not a function")}
+ return
+ }
+ result, err := fn(goja.Undefined(), p.vm.ToValue(req.url))
+ if err != nil {
+ req.respCh <- jsParserResp{err: err}
+ return
}
- if err := json.Unmarshal(data, &item); err != nil {
- req.respCh <- jsParserResp{err: fmt.Errorf("failed to unmarshal JSON to Item: %w", err)}
- continue
+ var item parser.Item
+ if exported := result.Export(); exported != nil {
+ data, err := json.Marshal(exported)
+ if err != nil {
+ req.respCh <- jsParserResp{err: fmt.Errorf("failed to marshal result to JSON: %w", err)}
+ return
+ }
+
+ if err := json.Unmarshal(data, &item); err != nil {
+ req.respCh <- jsParserResp{err: fmt.Errorf("failed to unmarshal JSON to Item: %w", err)}
+ return
+ }
+ } else {
+ req.respCh <- jsParserResp{err: fmt.Errorf("JS function returned null or undefined")}
+ return
}
- } else {
- req.respCh <- jsParserResp{err: fmt.Errorf("JS function returned null or undefined")}
- continue
+ req.respCh <- jsParserResp{item: &item}
}
- req.respCh <- jsParserResp{item: &item}
- }
+ }()
}
}()
@@ -116,7 +145,8 @@ func LoadPlugins(ctx context.Context, dir string) error {
scriptPath := filepath.Join(dir, e.Name())
code, err := os.ReadFile(scriptPath)
if err != nil {
- return err
+ log.Warnf("Failed to read plugin file %s: %v", e.Name(), err)
+ continue
}
vm := goja.New()
@@ -130,7 +160,8 @@ func LoadPlugins(ctx context.Context, dir string) error {
vm.Set("playwright", jsPlaywright(vm, logger))
if _, err := vm.RunString(string(code)); err != nil {
- return fmt.Errorf("error loading plugin %s: %w", e.Name(), err)
+ logger.Warnf("Failed to load plugin %s: %v", e.Name(), err)
+ continue
}
}
return nil
@@ -164,8 +195,12 @@ func addPlugin(ctx context.Context, code string, name string) error {
if len(configuredDirs) > 0 {
dir = configuredDirs[0]
}
+ fileName := filepath.Base(name)
+ if fileName == "" || fileName == "." || fileName == ".." {
+ return fmt.Errorf("invalid plugin name %q", name)
+ }
if err := os.MkdirAll(dir, 0755); err == nil {
- pluginPath := filepath.Join(dir, name)
+ pluginPath := filepath.Join(dir, fileName)
if err := os.WriteFile(pluginPath, []byte(code), 0644); err != nil {
logger.Warn("Failed to save plugin file: " + err.Error())
}
diff --git a/parsers/js/plugin_version_test.go b/parsers/js/plugin_version_test.go
new file mode 100644
index 0000000..bb4d682
--- /dev/null
+++ b/parsers/js/plugin_version_test.go
@@ -0,0 +1,35 @@
+package js_test
+
+import (
+ "os"
+ "path/filepath"
+ "testing"
+
+ "github.com/krau/SaveAny-Bot/parsers/js"
+)
+
+// Regression: a plugin with an invalid semver version must be rejected
+// without panicking the process (previously semver.MustParse crashed the bot).
+func TestLoadPluginsRejectsInvalidVersion(t *testing.T) {
+ dir := t.TempDir()
+ bad := `registerParser({
+ metadata: { name: "probe", version: "not-a-version", description: "", author: "" },
+ canHandle: function(url) { return true; },
+ parse: async function(url) { return { resources: [] }; }
+ });`
+ if err := os.WriteFile(filepath.Join(dir, "bad.js"), []byte(bad), 0o644); err != nil {
+ t.Fatal(err)
+ }
+ good := `registerParser({
+ metadata: { name: "good", version: "1.0.0", description: "", author: "" },
+ canHandle: function(url) { return false; },
+ parse: async function(url) { return { resources: [] }; }
+ });`
+ if err := os.WriteFile(filepath.Join(dir, "good.js"), []byte(good), 0o644); err != nil {
+ t.Fatal(err)
+ }
+ // Must not panic; the bad plugin is skipped and remaining plugins load.
+ if err := js.LoadPlugins(t.Context(), dir); err != nil {
+ t.Fatalf("LoadPlugins returned error: %v", err)
+ }
+}
diff --git a/parsers/native/kemono/kemono.go b/parsers/native/kemono/kemono.go
index 19f6172..243c097 100644
--- a/parsers/native/kemono/kemono.go
+++ b/parsers/native/kemono/kemono.go
@@ -10,6 +10,7 @@ import (
"path"
"strings"
+ "github.com/charmbracelet/log"
"github.com/duke-git/lancet/v2/strutil"
"github.com/krau/SaveAny-Bot/common/utils/netutil"
"github.com/krau/SaveAny-Bot/pkg/parser"
@@ -135,9 +136,17 @@ func (k *KemonoParser) parseOne(ctx context.Context, info *DownloadInfo) (*parse
if preview.Type == nil || *preview.Type != "thumbnail" {
continue
}
+ if preview.Path == nil || preview.Server == nil {
+ log.FromContext(ctx).Warnf("Skipping kemono preview with missing path or server: post %s", info.PostID)
+ continue
+ }
picCdnMap[*preview.Path] = *preview.Server
}
for _, attachment := range postInfo.Post.Attachments {
+ if attachment.Path == nil || attachment.Name == nil {
+ log.FromContext(ctx).Warnf("Skipping kemono post attachment with missing path or name: post %s", info.PostID)
+ continue
+ }
if !isImageExt(*attachment.Path) {
continue
}
diff --git a/parsers/native/kemono/post_legacy.go b/parsers/native/kemono/post_legacy.go
deleted file mode 100644
index 17146fc..0000000
--- a/parsers/native/kemono/post_legacy.go
+++ /dev/null
@@ -1,16 +0,0 @@
-package kemono
-
-type PostLegacy struct {
- Props Props `json:"props"`
- Results []Result `json:"results"`
-}
-
-type Props struct {
- Count uint `json:"count"`
- Limit uint `json:"limit"`
-}
-
-type Result struct {
- ID string `json:"id"`
- Title string `json:"title"`
-}
diff --git a/parsers/native/kemono/user_profile.go b/parsers/native/kemono/user_profile.go
deleted file mode 100644
index 8d7632e..0000000
--- a/parsers/native/kemono/user_profile.go
+++ /dev/null
@@ -1,8 +0,0 @@
-package kemono
-
-type UserProfile struct {
- ID string `json:"id"`
- Name string `json:"name"`
- Service string `json:"service"`
- PublicID *string `json:"public_id,omitempty"`
-}
diff --git a/parsers/native/twitter/types.go b/parsers/native/twitter/types.go
index 255a31e..98d1899 100644
--- a/parsers/native/twitter/types.go
+++ b/parsers/native/twitter/types.go
@@ -1,101 +1,5 @@
package twitter
-// type AutoGenerated struct {
-// Code int `json:"code"`
-// Message string `json:"message"`
-// Tweet struct {
-// URL string `json:"url"`
-// ID string `json:"id"`
-// Text string `json:"text"`
-// RawText struct {
-// Text string `json:"text"`
-// Facets []struct {
-// Type string `json:"type"`
-// Indices []int `json:"indices"`
-// Original string `json:"original"`
-// ID string `json:"id,omitempty"`
-// Display string `json:"display,omitempty"`
-// Replacement string `json:"replacement,omitempty"`
-// } `json:"facets"`
-// } `json:"raw_text"`
-// Author struct {
-// ID string `json:"id"`
-// Name string `json:"name"`
-// ScreenName string `json:"screen_name"`
-// AvatarURL string `json:"avatar_url"`
-// BannerURL interface{} `json:"banner_url"`
-// Description string `json:"description"`
-// Location string `json:"location"`
-// URL string `json:"url"`
-// Followers int `json:"followers"`
-// Following int `json:"following"`
-// Joined string `json:"joined"`
-// Likes int `json:"likes"`
-// MediaCount int `json:"media_count"`
-// Protected bool `json:"protected"`
-// Website struct {
-// URL string `json:"url"`
-// DisplayURL string `json:"display_url"`
-// } `json:"website"`
-// Tweets int `json:"tweets"`
-// AvatarColor interface{} `json:"avatar_color"`
-// } `json:"author"`
-// Replies int `json:"replies"`
-// Retweets int `json:"retweets"`
-// Likes int `json:"likes"`
-// Bookmarks int `json:"bookmarks"`
-// CreatedAt string `json:"created_at"`
-// CreatedTimestamp int `json:"created_timestamp"`
-// PossiblySensitive bool `json:"possibly_sensitive"`
-// Views int `json:"views"`
-// IsNoteTweet bool `json:"is_note_tweet"`
-// CommunityNote interface{} `json:"community_note"`
-// Lang string `json:"lang"`
-// ReplyingTo interface{} `json:"replying_to"`
-// ReplyingToStatus interface{} `json:"replying_to_status"`
-// Media struct {
-// All []struct {
-// URL string `json:"url"`
-// ThumbnailURL string `json:"thumbnail_url"`
-// Duration int `json:"duration"`
-// Width int `json:"width"`
-// Height int `json:"height"`
-// Format string `json:"format"`
-// Type string `json:"type"`
-// Variants []struct {
-// Bitrate int `json:"bitrate"`
-// ContentType string `json:"content_type"`
-// URL string `json:"url"`
-// } `json:"variants"`
-// } `json:"all"`
-// Photos []struct {
-// Type string `json:"type"`
-// URL string `json:"url"`
-// Width int `json:"width"`
-// Height int `json:"height"`
-// } `json:"photos"`
-// Videos []struct {
-// URL string `json:"url"`
-// ThumbnailURL string `json:"thumbnail_url"`
-// Duration int `json:"duration"`
-// Width int `json:"width"`
-// Height int `json:"height"`
-// Format string `json:"format"`
-// Type string `json:"type"`
-// Variants []struct {
-// Bitrate int `json:"bitrate"`
-// ContentType string `json:"content_type"`
-// URL string `json:"url"`
-// } `json:"variants"`
-// } `json:"videos"`
-// } `json:"media"`
-// Source string `json:"source"`
-// TwitterCard string `json:"twitter_card"`
-// Color interface{} `json:"color"`
-// Provider string `json:"provider"`
-// } `json:"tweet"`
-// }
-
type FxTwitterApiResp struct {
Code int `json:"code"`
Message string `json:"message"`
diff --git a/parsers/parsers/parsers.go b/parsers/parsers/parsers.go
index ebfd483..b3c15a3 100644
--- a/parsers/parsers/parsers.go
+++ b/parsers/parsers/parsers.go
@@ -2,6 +2,7 @@ package parsers
import (
"fmt"
+ "slices"
"sync"
"github.com/krau/SaveAny-Bot/config"
@@ -39,5 +40,5 @@ func Get() []parser.Parser {
configOnce.Do(configParsers)
mu.Lock()
defer mu.Unlock()
- return parsers
+ return slices.Clone(parsers)
}
diff --git a/pkg/consts/tglimit/tglimit.go b/pkg/consts/tglimit/tglimit.go
index 94c89db..7bd49f5 100644
--- a/pkg/consts/tglimit/tglimit.go
+++ b/pkg/consts/tglimit/tglimit.go
@@ -8,4 +8,7 @@ const (
MaxPartSize = 1024 * 1024
MaxUploadPartSize = uploader.MaximumPartSize
MaxPhotoSize = 10 * 1024 * 1024
+ // MaxAlbumItems is the Telegram media-album item cap used for batching
+ // uploads and lossless video splitting.
+ MaxAlbumItems = 10
)
diff --git a/pkg/parser/parser.go b/pkg/parser/parser.go
index b537df5..3044eac 100644
--- a/pkg/parser/parser.go
+++ b/pkg/parser/parser.go
@@ -4,6 +4,8 @@ import (
"context"
"crypto/md5"
"fmt"
+ "maps"
+ "slices"
)
type Parser interface {
@@ -57,14 +59,15 @@ func (r *Resource) ID() string {
h.Write([]byte(r.Extension))
fmt.Fprintf(h, "%d", r.Size)
- for k, v := range r.Hash {
+ // Sort keys for a stable fingerprint.
+ for _, k := range slices.Sorted(maps.Keys(r.Hash)) {
h.Write([]byte(k))
- h.Write([]byte(v))
+ h.Write([]byte(r.Hash[k]))
}
- for k, v := range r.Headers {
+ for _, k := range slices.Sorted(maps.Keys(r.Headers)) {
h.Write([]byte(k))
- h.Write([]byte(v))
+ h.Write([]byte(r.Headers[k]))
}
return fmt.Sprintf("%x", h.Sum(nil))
diff --git a/pkg/queue/queue.go b/pkg/queue/queue.go
index 145e174..47c248d 100644
--- a/pkg/queue/queue.go
+++ b/pkg/queue/queue.go
@@ -50,38 +50,41 @@ func (tq *TaskQueue[T]) Add(task *Task[T]) error {
return nil
}
+// ErrQueueClosed is returned by Get when the queue is closed and no tasks remain.
+var ErrQueueClosed = errors.New("queue is closed and empty")
+
// Get retrieves and removes the next non-cancelled task from the queue, adding it to the running tasks.
// Blocks until a task is available or the queue is closed.
func (tq *TaskQueue[T]) Get() (*Task[T], error) {
tq.mu.Lock()
defer tq.mu.Unlock()
- for tq.tasks.Len() == 0 && !tq.closed {
- tq.cond.Wait()
- }
+ for {
+ for tq.tasks.Len() == 0 && !tq.closed {
+ tq.cond.Wait()
+ }
- if tq.closed && tq.tasks.Len() == 0 {
- return nil, fmt.Errorf("queue is closed and empty")
- }
+ for tq.tasks.Len() > 0 {
+ element := tq.tasks.Front()
+ task := element.Value.(*Task[T])
- for tq.tasks.Len() > 0 {
- element := tq.tasks.Front()
- task := element.Value.(*Task[T])
+ tq.tasks.Remove(element)
+ task.element = nil
- tq.tasks.Remove(element)
- task.element = nil
+ if task.Cancelled() {
+ // Skip cancelled tasks and release their IDs.
+ delete(tq.taskMap, task.ID)
+ continue
+ }
- if !task.Cancelled() {
tq.runningTaskMap[task.ID] = task
return task, nil
}
- }
- if !tq.closed {
- return tq.Get()
+ if tq.closed {
+ return nil, ErrQueueClosed
+ }
}
-
- return nil, fmt.Errorf("queue is closed and empty")
}
// Done stops(cancels) and removes the task from the running tasks.
diff --git a/pkg/queue/queue_test.go b/pkg/queue/queue_test.go
index b36b1c8..f6bba20 100644
--- a/pkg/queue/queue_test.go
+++ b/pkg/queue/queue_test.go
@@ -2,9 +2,11 @@ package queue_test
import (
"context"
+ "errors"
"fmt"
"sync"
"testing"
+ "time"
"github.com/krau/SaveAny-Bot/pkg/queue"
)
@@ -65,8 +67,8 @@ func TestCloseBehavior(t *testing.T) {
// consumer
go func() {
_, err := q.Get()
- if err == nil {
- t.Errorf("expected error when getting from closed empty queue, got nil")
+ if !errors.Is(err, queue.ErrQueueClosed) {
+ t.Errorf("expected ErrQueueClosed from closed empty queue, got %v", err)
}
close(done)
}()
@@ -77,6 +79,86 @@ func TestCloseBehavior(t *testing.T) {
<-done
}
+// Regression: Get() must not deadlock when every queued task was cancelled
+// before a worker picked it up (previously recursed while holding the mutex).
+func TestGetAfterAllCancelled(t *testing.T) {
+ q := queue.NewTaskQueue[int]()
+ for i := range 3 {
+ if err := q.Add(newTask(fmt.Sprintf("c%d", i))); err != nil {
+ t.Fatalf("unexpected error on Add: %v", err)
+ }
+ }
+ for i := range 3 {
+ if err := q.CancelTask(fmt.Sprintf("c%d", i)); err != nil {
+ t.Fatalf("unexpected error on CancelTask: %v", err)
+ }
+ }
+
+ // A task added after the cancelled ones must still be delivered.
+ if err := q.Add(newTask("late")); err != nil {
+ t.Fatalf("unexpected error on Add after cancel: %v", err)
+ }
+
+ done := make(chan struct{})
+ var got *queue.Task[int]
+ go func() {
+ var err error
+ got, err = q.Get()
+ if err != nil {
+ t.Errorf("unexpected error on Get: %v", err)
+ }
+ close(done)
+ }()
+ select {
+ case <-done:
+ if got == nil || got.ID != "late" {
+ t.Fatalf("expected task 'late', got %v", got)
+ }
+ case <-time.After(2 * time.Second):
+ t.Fatal("Get() deadlocked after all queued tasks were cancelled")
+ }
+}
+
+// Cancelled queued tasks must be dropped from the task map so their TaskID
+// can be reused (previously they leaked until process exit).
+func TestCancelledQueuedTaskIDReusable(t *testing.T) {
+ q := queue.NewTaskQueue[int]()
+ if err := q.Add(newTask("dup")); err != nil {
+ t.Fatal(err)
+ }
+ if err := q.CancelTask("dup"); err != nil {
+ t.Fatal(err)
+ }
+ done := make(chan struct{})
+ var gotID string
+ go func() {
+ task, err := q.Get()
+ if err != nil {
+ t.Errorf("unexpected error on Get: %v", err)
+ close(done)
+ return
+ }
+ gotID = task.ID
+ close(done)
+ }()
+ // The first Get skips the cancelled task and blocks; a live task unblocks it.
+ if err := q.Add(newTask("late")); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case <-done:
+ if gotID != "late" {
+ t.Fatalf("expected 'late', got %q", gotID)
+ }
+ case <-time.After(2 * time.Second):
+ t.Fatal("Get did not return the live task")
+ }
+ // The cancelled task was dropped from the map: its ID is reusable.
+ if err := q.Add(newTask("dup")); err != nil {
+ t.Fatalf("expected cancelled task ID to be reusable, got: %v", err)
+ }
+}
+
func TestConcurrencySafety(t *testing.T) {
q := queue.NewTaskQueue[int]()
var wg sync.WaitGroup
diff --git a/storage/alist/alist.go b/storage/alist/alist.go
index cc0a2ee..8afc0a2 100644
--- a/storage/alist/alist.go
+++ b/storage/alist/alist.go
@@ -9,23 +9,35 @@ import (
"net/http"
"net/url"
"path"
- "strings"
+ "sync"
"time"
"github.com/charmbracelet/log"
+ "github.com/krau/SaveAny-Bot/common/utils/fsutil"
config "github.com/krau/SaveAny-Bot/config/storage"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
"github.com/krau/SaveAny-Bot/pkg/storagetypes"
+ "golang.org/x/sync/singleflight"
)
type Alist struct {
- client *http.Client
- token string
- baseURL string
- loginInfo *loginRequest
- config config.AlistStorageConfig
- logger *log.Logger
+ client *http.Client
+ tokenMu sync.RWMutex
+ token string
+ lastLoginAt time.Time
+ tokenFlight singleflight.Group
+ baseURL string
+ loginInfo *loginRequest
+ config config.AlistStorageConfig
+ logger *log.Logger
+}
+
+// authHeader returns the current token for use in API requests.
+func (a *Alist) authHeader() string {
+ a.tokenMu.RLock()
+ defer a.tokenMu.RUnlock()
+ return a.token
}
func (a *Alist) Init(ctx context.Context, cfg config.StorageConfig) error {
@@ -42,39 +54,35 @@ func (a *Alist) Init(ctx context.Context, cfg config.StorageConfig) error {
a.logger = log.FromContext(ctx).WithPrefix(fmt.Sprintf("alist[%s]", alistConfig.Name))
if alistConfig.Token != "" {
+ a.tokenMu.Lock()
a.token = alistConfig.Token
- ctx, cancel := context.WithTimeout(context.Background(), 1*time.Minute)
+ a.tokenMu.Unlock()
+ tokenCtx, cancel := context.WithTimeout(ctx, 1*time.Minute)
defer cancel()
- req, err := http.NewRequestWithContext(ctx, http.MethodGet, a.baseURL+"/api/me", nil)
+ req, err := http.NewRequestWithContext(tokenCtx, http.MethodGet, a.baseURL+"/api/me", nil)
if err != nil {
- a.logger.Fatalf("Failed to create request: %v", err)
- return err
+ return fmt.Errorf("failed to create request: %w", err)
}
- req.Header.Set("Authorization", a.token)
+ req.Header.Set("Authorization", a.authHeader())
resp, err := a.client.Do(req)
if err != nil {
- a.logger.Fatalf("Failed to send request: %v", err)
- return err
+ return fmt.Errorf("failed to send request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
- a.logger.Fatalf("Failed to get alist user info: %s", resp.Status)
- return err
+ return fmt.Errorf("failed to get alist user info: %s", resp.Status)
}
body, err := io.ReadAll(resp.Body)
if err != nil {
- a.logger.Fatalf("Failed to read response body: %v", err)
- return err
+ return fmt.Errorf("failed to read response body: %w", err)
}
var meResp meResponse
if err := json.Unmarshal(body, &meResp); err != nil {
- a.logger.Fatalf("Failed to unmarshal me response: %v", err)
- return err
+ return fmt.Errorf("failed to unmarshal me response: %w", err)
}
if meResp.Code != http.StatusOK {
- a.logger.Fatalf("Failed to get alist user info: %s", meResp.Message)
- return err
+ return fmt.Errorf("failed to get alist user info: %s", meResp.Message)
}
a.logger.Debugf("Logged in Alist as %s", meResp.Data.Username)
return nil
@@ -85,12 +93,15 @@ func (a *Alist) Init(ctx context.Context, cfg config.StorageConfig) error {
}
if err := a.getToken(ctx); err != nil {
- a.logger.Fatalf("Failed to login to Alist: %v", err)
- return err
+ return fmt.Errorf("failed to login to Alist: %w", err)
}
+ // The init login must not satisfy the refresh dedup window.
+ a.tokenMu.Lock()
+ a.lastLoginAt = time.Time{}
+ a.tokenMu.Unlock()
a.logger.Debug("Logged in to Alist")
- go a.refreshToken(*alistConfig)
+ go a.refreshToken(ctx, *alistConfig)
return nil
}
@@ -104,33 +115,40 @@ func (a *Alist) Name() string {
func (a *Alist) Save(ctx context.Context, reader io.Reader, storagePath string) error {
a.logger.Infof("Saving file to %s", storagePath)
- storagePath = a.JoinStoragePath(storagePath)
- ext := path.Ext(storagePath)
- base := strings.TrimSuffix(storagePath, ext)
- candidate := storagePath
+ candidate := a.JoinStoragePath(storagePath)
if overwrite, _ := ctx.Value(ctxkey.OverwriteExisting).(bool); !overwrite {
- for i := 1; a.existsPath(ctx, candidate); i++ {
- candidate = fmt.Sprintf("%s_%d%s", base, i, ext)
- }
+ candidate = fsutil.UniquePath(a.config.BasePath, storagePath, func(c string) bool {
+ return a.existsPath(ctx, c)
+ }, 1000)
}
- req, err := http.NewRequestWithContext(ctx, http.MethodPut, a.baseURL+"/api/fs/put", reader)
+ resp, err := a.putFile(ctx, reader, candidate)
if err != nil {
- return fmt.Errorf("failed to create request: %w", err)
+ return err
}
- req.Header.Set("Authorization", a.token)
- req.Header.Set("File-Path", url.PathEscape(candidate))
- req.Header.Set("Content-Type", "application/octet-stream")
- if length := ctx.Value(ctxkey.ContentLength); length != nil {
- length, ok := length.(int64)
- if ok {
- req.ContentLength = length
+ if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
+ status := resp.Status
+ resp.Body.Close()
+ // Token-only storage cannot refresh: surface the auth error.
+ if a.loginInfo == nil {
+ return fmt.Errorf("failed to save file to Alist: %s", status)
+ }
+ if err := a.getToken(ctx); err != nil {
+ return fmt.Errorf("failed to refresh alist token: %w", err)
+ }
+ rs, seekable := reader.(io.ReadSeeker)
+ if !seekable {
+ a.logger.Warnf("Upload rejected with %s; reader is not seekable, cannot retry", status)
+ return fmt.Errorf("failed to save file to Alist: %s (streaming reader cannot be replayed for retry)", status)
+ }
+ if _, err := rs.Seek(0, io.SeekStart); err != nil {
+ return fmt.Errorf("failed to rewind reader before retry: %w", err)
+ }
+ a.logger.Info("Retrying upload with refreshed token")
+ resp, err = a.putFile(ctx, reader, candidate)
+ if err != nil {
+ return err
}
- }
-
- resp, err := a.client.Do(req)
- if err != nil {
- return fmt.Errorf("failed to send request: %w", err)
}
defer resp.Body.Close()
@@ -155,6 +173,29 @@ func (a *Alist) Save(ctx context.Context, reader io.Reader, storagePath string)
return nil
}
+// putFile performs a single PUT upload with the given token and returns the response.
+func (a *Alist) putFile(ctx context.Context, reader io.Reader, storagePath string) (*http.Response, error) {
+ req, err := http.NewRequestWithContext(ctx, http.MethodPut, a.baseURL+"/api/fs/put", reader)
+ if err != nil {
+ return nil, fmt.Errorf("failed to create request: %w", err)
+ }
+ req.Header.Set("Authorization", a.authHeader())
+ req.Header.Set("File-Path", url.PathEscape(storagePath))
+ req.Header.Set("Content-Type", "application/octet-stream")
+ if length := ctx.Value(ctxkey.ContentLength); length != nil {
+ length, ok := length.(int64)
+ if ok {
+ req.ContentLength = length
+ }
+ }
+
+ resp, err := a.client.Do(req)
+ if err != nil {
+ return nil, fmt.Errorf("failed to send request: %w", err)
+ }
+ return resp, nil
+}
+
func (a *Alist) JoinStoragePath(p string) string {
return path.Join(a.config.BasePath, p)
}
@@ -189,7 +230,7 @@ func (a *Alist) existsPath(ctx context.Context, storagePath string) bool {
a.logger.Errorf("Failed to create request: %v", err)
return false
}
- req.Header.Set("Authorization", a.token)
+ req.Header.Set("Authorization", a.authHeader())
req.Header.Set("Content-Type", "application/json")
resp, err := a.client.Do(req)
if err != nil {
@@ -244,7 +285,7 @@ func (a *Alist) ListFiles(ctx context.Context, dirPath string) ([]storagetypes.F
if err != nil {
return nil, fmt.Errorf("failed to create request: %w", err)
}
- req.Header.Set("Authorization", a.token)
+ req.Header.Set("Authorization", a.authHeader())
req.Header.Set("Content-Type", "application/json")
resp, err := a.client.Do(req)
@@ -319,7 +360,7 @@ func (a *Alist) OpenFile(ctx context.Context, filePath string) (io.ReadCloser, i
if err != nil {
return nil, 0, fmt.Errorf("failed to create request: %w", err)
}
- req.Header.Set("Authorization", a.token)
+ req.Header.Set("Authorization", a.authHeader())
req.Header.Set("Content-Type", "application/json")
resp, err := a.client.Do(req)
diff --git a/storage/alist/alist_test.go b/storage/alist/alist_test.go
new file mode 100644
index 0000000..03175fa
--- /dev/null
+++ b/storage/alist/alist_test.go
@@ -0,0 +1,152 @@
+package alist_test
+
+import (
+ "bytes"
+ "encoding/json"
+ "fmt"
+ "net/http"
+ "net/http/httptest"
+ "strings"
+ "sync"
+ "testing"
+
+ storconfig "github.com/krau/SaveAny-Bot/config/storage"
+ "github.com/krau/SaveAny-Bot/storage/alist"
+)
+
+// newAlistServer starts a fake alist whose login endpoint issues sequential
+// tokens and whose PUT endpoint rejects the given token (simulating an expired
+// credential) while accepting refreshed ones.
+func newAlistServer(t *testing.T, rejectedToken string) (*httptest.Server, *sync.Mutex, *int, *[]putRecord) {
+ t.Helper()
+ var mu sync.Mutex
+ loginCount := 0
+ tokenSeq := 0
+ var puts []putRecord
+
+ mux := http.NewServeMux()
+ mux.HandleFunc("/api/auth/login", func(w http.ResponseWriter, r *http.Request) {
+ mu.Lock()
+ loginCount++
+ tokenSeq++
+ token := fmt.Sprintf("token-%d", tokenSeq)
+ mu.Unlock()
+ w.Header().Set("Content-Type", "application/json")
+ json.NewEncoder(w).Encode(map[string]any{"code": 200, "message": "ok", "data": map[string]any{"token": token}})
+ })
+ mux.HandleFunc("/api/me", func(w http.ResponseWriter, r *http.Request) {
+ w.Header().Set("Content-Type", "application/json")
+ json.NewEncoder(w).Encode(map[string]any{"code": 200, "message": "ok", "data": map[string]any{"username": "probe"}})
+ })
+ mux.HandleFunc("/api/fs/put", func(w http.ResponseWriter, r *http.Request) {
+ mu.Lock()
+ rejected := r.Header.Get("Authorization") == rejectedToken
+ puts = append(puts, putRecord{auth: r.Header.Get("Authorization"), rejected: rejected})
+ mu.Unlock()
+ w.Header().Set("Content-Type", "application/json")
+ if rejected {
+ w.WriteHeader(http.StatusUnauthorized)
+ json.NewEncoder(w).Encode(map[string]any{"code": 401, "message": "unauthorized"})
+ return
+ }
+ json.NewEncoder(w).Encode(map[string]any{"code": 200, "message": "ok"})
+ })
+ srv := httptest.NewServer(mux)
+ t.Cleanup(srv.Close)
+ return srv, &mu, &loginCount, &puts
+}
+
+type putRecord struct {
+ auth string
+ rejected bool
+}
+
+// Regression: concurrent uploads hitting 401 must share a single re-login
+// (singleflight) and retry with the refreshed token. Init performs login #1
+// (token-1); the server rejects it, so the concurrent uploads must trigger a
+// second, merged login (token-2).
+func TestConcurrent401RetrySingleLogin(t *testing.T) {
+ srv, mu, loginCount, putAuths := newAlistServer(t, "token-1")
+
+ cfg := &storconfig.AlistStorageConfig{}
+ cfg.Name = "probe"
+ cfg.URL = srv.URL
+ cfg.Username = "user"
+ cfg.Password = "pass"
+ cfg.BasePath = "/probe"
+
+ stor := &alist.Alist{}
+ if err := stor.Init(t.Context(), cfg); err != nil {
+ t.Fatalf("Init failed: %v", err)
+ }
+
+ const workers = 10
+ var wg sync.WaitGroup
+ errs := make(chan error, workers)
+ for range workers {
+ wg.Go(func() {
+ errs <- stor.Save(t.Context(), bytes.NewReader([]byte("data")), "dir/file.txt")
+ })
+ }
+ wg.Wait()
+ close(errs)
+ for err := range errs {
+ if err != nil {
+ t.Fatalf("Save failed: %v", err)
+ }
+ }
+
+ mu.Lock()
+ defer mu.Unlock()
+ // One login for init, one merged login for the 401 storm.
+ if *loginCount != 2 {
+ t.Fatalf("expected 2 logins (init + merged retry), got %d", *loginCount)
+ }
+ accepted := 0
+ for _, put := range *putAuths {
+ if put.rejected {
+ if put.auth != "token-1" {
+ t.Fatalf("expected rejected uploads to use the expired token-1, got %q", put.auth)
+ }
+ continue
+ }
+ accepted++
+ if put.auth != "token-2" {
+ t.Fatalf("expected accepted uploads to use the refreshed token-2, got %q", put.auth)
+ }
+ }
+ if accepted != workers {
+ t.Fatalf("expected %d accepted uploads, got %d", workers, accepted)
+ }
+}
+
+// A token-only storage receives 401 and must return the auth error without
+// attempting a login (it has no credentials to refresh with).
+func TestTokenOnlyNoLoginOn401(t *testing.T) {
+ srv, mu, loginCount, _ := newAlistServer(t, "token-0")
+
+ cfg := &storconfig.AlistStorageConfig{}
+ cfg.Name = "probe"
+ cfg.URL = srv.URL
+ cfg.Token = "token-0"
+ cfg.BasePath = "/probe"
+
+ stor := &alist.Alist{}
+ if err := stor.Init(t.Context(), cfg); err != nil {
+ t.Fatalf("Init failed: %v", err)
+ }
+
+ err := stor.Save(t.Context(), bytes.NewReader([]byte("data")), "dir/file.txt")
+ if err == nil {
+ t.Fatal("expected auth error from token-only storage, got nil")
+ }
+ if !strings.Contains(err.Error(), "401") {
+ t.Fatalf("expected 401 auth error, got: %v", err)
+ }
+
+ mu.Lock()
+ defer mu.Unlock()
+ if *loginCount != 0 {
+ t.Fatalf("expected no login attempts for token-only storage, got %d", *loginCount)
+ }
+}
diff --git a/storage/alist/token.go b/storage/alist/token.go
index e2bc10e..fa0e8fe 100644
--- a/storage/alist/token.go
+++ b/storage/alist/token.go
@@ -12,7 +12,42 @@ import (
config "github.com/krau/SaveAny-Bot/config/storage"
)
+// minTokenRefreshInterval deduplicates login storms; it is capped by the
+// configured token expiry.
+const minTokenRefreshInterval = 30 * time.Second
+
+func (a *Alist) tokenRefreshWindow() time.Duration {
+ window := minTokenRefreshInterval
+ if exp := time.Duration(a.config.TokenExp) * time.Second; exp > 0 && exp < window {
+ window = exp
+ }
+ return window
+}
+
+// getToken refreshes the JWT, merging concurrent calls.
func (a *Alist) getToken(ctx context.Context) error {
+ a.tokenMu.RLock()
+ fresh := !a.lastLoginAt.IsZero() && time.Since(a.lastLoginAt) < a.tokenRefreshWindow()
+ a.tokenMu.RUnlock()
+ if fresh {
+ return nil
+ }
+ _, err, _ := a.tokenFlight.Do("token", func() (any, error) {
+ a.tokenMu.RLock()
+ fresh := !a.lastLoginAt.IsZero() && time.Since(a.lastLoginAt) < a.tokenRefreshWindow()
+ a.tokenMu.RUnlock()
+ if fresh {
+ return nil, nil
+ }
+ return nil, a.fetchToken(ctx)
+ })
+ return err
+}
+
+func (a *Alist) fetchToken(ctx context.Context) error {
+ if a.loginInfo == nil {
+ return fmt.Errorf("token-only alist storage cannot refresh credentials")
+ }
loginBody, err := json.Marshal(a.loginInfo)
if err != nil {
return fmt.Errorf("failed to marshal login request: %w", err)
@@ -44,22 +79,31 @@ func (a *Alist) getToken(ctx context.Context) error {
return fmt.Errorf("%w: %s", ErrAlistLoginFailed, loginResp.Message)
}
+ a.tokenMu.Lock()
a.token = loginResp.Data.Token
+ a.lastLoginAt = time.Now()
+ a.tokenMu.Unlock()
return nil
}
-func (a *Alist) refreshToken(cfg config.AlistStorageConfig) {
+func (a *Alist) refreshToken(ctx context.Context, cfg config.AlistStorageConfig) {
tokenExp := cfg.TokenExp
if tokenExp <= 0 {
a.logger.Warn("Invalid token expiration time, using default value")
tokenExp = 3600
}
+ ticker := time.NewTicker(time.Duration(tokenExp) * time.Second)
+ defer ticker.Stop()
for {
- time.Sleep(time.Duration(tokenExp) * time.Second)
- if err := a.getToken(context.Background()); err != nil {
- a.logger.Errorf("Failed to refresh jwt token: %v", err)
- continue
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ if err := a.getToken(ctx); err != nil {
+ a.logger.Errorf("Failed to refresh jwt token: %v", err)
+ continue
+ }
+ a.logger.Info("Refreshed Alist jwt token")
}
- a.logger.Info("Refreshed Alist jwt token")
}
}
diff --git a/storage/load.go b/storage/load.go
index 09bf4bf..67b916d 100644
--- a/storage/load.go
+++ b/storage/load.go
@@ -3,13 +3,44 @@ package storage
import (
"context"
"fmt"
+ "maps"
+ "sync"
"github.com/charmbracelet/log"
"github.com/krau/SaveAny-Bot/config"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
+ "golang.org/x/sync/singleflight"
)
-var UserStorages = make(map[int64][]Storage)
+var (
+ storageMu sync.RWMutex
+ // Storages maps storage names to initialized storage instances.
+ Storages = make(map[string]Storage)
+
+ userStoragesMu sync.RWMutex
+ // UserStorages maps user IDs to their available storage instances.
+ UserStorages = make(map[int64][]Storage)
+
+ initFlight singleflight.Group
+)
+
+// GetStorage returns the initialized storage instance for name, without
+// creating one on demand.
+func GetStorage(name string) (Storage, bool) {
+ storageMu.RLock()
+ defer storageMu.RUnlock()
+ s, ok := Storages[name]
+ return s, ok
+}
+
+// AllStorages returns a snapshot copy of all initialized storages.
+func AllStorages() map[string]Storage {
+ storageMu.RLock()
+ defer storageMu.RUnlock()
+ out := make(map[string]Storage, len(Storages))
+ maps.Copy(out, Storages)
+ return out
+}
// GetStorageByName returns storage by name from cache or creates new one
// It should NOT be used to get storage for user, use GetStorageByUserIDAndName instead
@@ -18,21 +49,41 @@ func GetStorageByName(ctx context.Context, name string) (Storage, error) {
return nil, ErrStorageNameEmpty
}
+ storageMu.RLock()
storage, ok := Storages[name]
+ storageMu.RUnlock()
if ok {
return storage, nil
}
cfg := config.C().GetStorageByName(name)
if cfg == nil {
- return nil, fmt.Errorf("未找到存储 %s", name)
+ return nil, fmt.Errorf("storage %s not found", name)
}
- storage, err := NewStorage(ctx, cfg)
+ // Merge concurrent first-time initializations.
+ v, err, _ := initFlight.Do("storage:"+name, func() (any, error) {
+ storageMu.RLock()
+ if existing, ok := Storages[name]; ok {
+ storageMu.RUnlock()
+ return existing, nil
+ }
+ storageMu.RUnlock()
+ storage, err := NewStorage(ctx, cfg)
+ if err != nil {
+ return nil, err
+ }
+ storageMu.Lock()
+ defer storageMu.Unlock()
+ if existing, ok := Storages[name]; ok {
+ return existing, nil
+ }
+ Storages[name] = storage
+ return storage, nil
+ })
if err != nil {
return nil, err
}
- Storages[name] = storage
- return storage, nil
+ return v.(Storage), nil
}
// 检查 user 是否可用指定的 storage, 若不可用则返回未找到错误
@@ -52,8 +103,11 @@ func GetUserStorages(ctx context.Context, chatID int64) []Storage {
if chatID <= 0 {
return nil
}
- if storages, ok := UserStorages[chatID]; ok {
- return storages
+ userStoragesMu.RLock()
+ cached, ok := UserStorages[chatID]
+ userStoragesMu.RUnlock()
+ if ok {
+ return cached
}
var storages []Storage
for _, name := range config.C().GetStorageNamesByUserID(chatID) {
@@ -75,9 +129,16 @@ func LoadStorages(ctx context.Context) {
logger.Errorf("failed to load storage %s: %v", storage.GetName(), err)
}
}
- logger.Infof("successfully loaded %d storages", len(Storages))
+ storageMu.RLock()
+ loaded := len(Storages)
+ storageMu.RUnlock()
+ logger.Infof("successfully loaded %d storages", loaded)
for user := range config.C().GetUsersID() {
- UserStorages[int64(user)] = GetUserStorages(ctx, int64(user))
+ uid := int64(user)
+ storages := GetUserStorages(ctx, uid)
+ userStoragesMu.Lock()
+ UserStorages[uid] = storages
+ userStoragesMu.Unlock()
}
}
diff --git a/storage/local/local.go b/storage/local/local.go
index 4031d59..1c7b4b6 100644
--- a/storage/local/local.go
+++ b/storage/local/local.go
@@ -10,6 +10,7 @@ import (
"github.com/charmbracelet/log"
"github.com/duke-git/lancet/v2/fileutil"
+ "github.com/krau/SaveAny-Bot/common/utils/fsutil"
config "github.com/krau/SaveAny-Bot/config/storage"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
@@ -52,15 +53,17 @@ func (l *Local) JoinStoragePath(path string) string {
func (l *Local) Save(ctx context.Context, r io.Reader, storagePath string) error {
l.logger.Infof("Saving file to %s", storagePath)
- storagePath = l.JoinStoragePath(storagePath)
+ storagePath = filepath.Clean(storagePath)
+ if filepath.IsAbs(storagePath) {
+ return fmt.Errorf("local: storage path must be relative: %s", storagePath)
+ }
+ if storagePath == ".." || strings.HasPrefix(storagePath, ".."+string(filepath.Separator)) {
+ return fmt.Errorf("local: storage path escapes base directory: %s", storagePath)
+ }
- ext := filepath.Ext(storagePath)
- base := strings.TrimSuffix(storagePath, ext)
- candidate := storagePath
+ candidate := l.JoinStoragePath(storagePath)
if overwrite, _ := ctx.Value(ctxkey.OverwriteExisting).(bool); !overwrite {
- for i := 1; l.existsPath(candidate); i++ {
- candidate = fmt.Sprintf("%s_%d%s", base, i, ext)
- }
+ candidate = fsutil.UniquePath(l.config.BasePath, storagePath, l.existsPath, 1000)
}
absPath, err := filepath.Abs(candidate)
@@ -68,13 +71,17 @@ func (l *Local) Save(ctx context.Context, r io.Reader, storagePath string) error
return err
}
if err := fileutil.CreateDir(filepath.Dir(absPath)); err != nil {
- return err
+ return fmt.Errorf("failed to create directory: %w", err)
}
file, err := os.Create(absPath)
if err != nil {
- return err
+ return fmt.Errorf("failed to create file: %w", err)
}
- defer file.Close()
+ defer func() {
+ if err := file.Close(); err != nil {
+ l.logger.Errorf("Failed to close file %s: %v", absPath, err)
+ }
+ }()
_, err = io.Copy(file, r)
return err
}
diff --git a/storage/minio/client.go b/storage/minio/client.go
index d62d031..a807868 100644
--- a/storage/minio/client.go
+++ b/storage/minio/client.go
@@ -11,12 +11,12 @@ import (
"sync"
"github.com/charmbracelet/log"
+ "github.com/krau/SaveAny-Bot/common/utils/fsutil"
config "github.com/krau/SaveAny-Bot/config/storage"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
"github.com/minio/minio-go/v7"
"github.com/minio/minio-go/v7/pkg/credentials"
- "github.com/rs/xid"
)
var (
@@ -77,19 +77,11 @@ func (m *Minio) JoinStoragePath(p string) string {
func (m *Minio) Save(ctx context.Context, r io.Reader, storagePath string) error {
m.logger.Infof("Saving file from reader to %s", storagePath)
- storagePath = m.JoinStoragePath(storagePath)
- ext := path.Ext(storagePath)
- base := strings.TrimSuffix(storagePath, ext)
- candidate := storagePath
+ candidate := m.JoinStoragePath(storagePath)
if overwrite, _ := ctx.Value(ctxkey.OverwriteExisting).(bool); !overwrite {
- for i := 1; m.existsObject(ctx, candidate); i++ {
- candidate = fmt.Sprintf("%s_%d%s", base, i, ext)
- if i > 10 {
- m.logger.Errorf("Too many attempts to find a unique filename for %s", storagePath)
- candidate = fmt.Sprintf("%s_%s%s", base, xid.New().String(), ext)
- break
- }
- }
+ candidate = fsutil.UniquePath(strings.TrimPrefix(m.config.BasePath, "/"), storagePath, func(c string) bool {
+ return m.existsObject(ctx, c)
+ }, 10)
}
size := int64(-1)
if length := ctx.Value(ctxkey.ContentLength); length != nil {
diff --git a/storage/rclone/errs.go b/storage/rclone/errs.go
index c430d3a..08f8407 100644
--- a/storage/rclone/errs.go
+++ b/storage/rclone/errs.go
@@ -8,7 +8,4 @@ var (
ErrFailedToSaveFile = errors.New("rclone: failed to save file")
ErrFailedToListFiles = errors.New("rclone: failed to list files")
ErrFailedToOpenFile = errors.New("rclone: failed to open file")
- ErrFailedToCheckFile = errors.New("rclone: failed to check file exists")
- ErrFailedToCreateDir = errors.New("rclone: failed to create directory")
- ErrCommandFailed = errors.New("rclone: command execution failed")
)
diff --git a/storage/rclone/rclone.go b/storage/rclone/rclone.go
index 9b2a0be..937d887 100644
--- a/storage/rclone/rclone.go
+++ b/storage/rclone/rclone.go
@@ -13,11 +13,11 @@ import (
"time"
"github.com/charmbracelet/log"
+ "github.com/krau/SaveAny-Bot/common/utils/fsutil"
config "github.com/krau/SaveAny-Bot/config/storage"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
"github.com/krau/SaveAny-Bot/pkg/storagetypes"
- "github.com/rs/xid"
)
type Rclone struct {
@@ -51,9 +51,6 @@ func (r *Rclone) Init(ctx context.Context, cfg config.StorageConfig) error {
}
remoteName := strings.TrimSuffix(r.config.Remote, ":")
- if !strings.HasSuffix(r.config.Remote, ":") {
- remoteName = r.config.Remote
- }
found := false
scanner := bufio.NewScanner(bytes.NewReader(output))
@@ -105,18 +102,11 @@ func (r *Rclone) getRemotePath(storagePath string) string {
func (r *Rclone) Save(ctx context.Context, reader io.Reader, storagePath string) error {
r.logger.Infof("Saving file to %s", storagePath)
- ext := path.Ext(storagePath)
- base := strings.TrimSuffix(storagePath, ext)
candidate := storagePath
if overwrite, _ := ctx.Value(ctxkey.OverwriteExisting).(bool); !overwrite {
- for i := 1; r.Exists(ctx, candidate); i++ {
- candidate = fmt.Sprintf("%s_%d%s", base, i, ext)
- if i > 100 {
- r.logger.Errorf("Too many attempts to find a unique filename for %s", storagePath)
- candidate = fmt.Sprintf("%s_%s%s", base, xid.New().String(), ext)
- break
- }
- }
+ candidate = fsutil.UniquePath("", storagePath, func(c string) bool {
+ return r.Exists(ctx, c)
+ }, 100)
}
remotePath := r.getRemotePath(candidate)
@@ -285,6 +275,12 @@ func (r *rcloneCatReader) Close() error {
if err := r.reader.Close(); err != nil {
r.logger.Warnf("Failed to close reader: %v", err)
}
+ // Kill the cat process so Wait cannot block on a hung pipe.
+ if r.cmd.Process != nil {
+ if err := r.cmd.Process.Kill(); err != nil {
+ r.logger.Warnf("Failed to kill rclone cat process: %v", err)
+ }
+ }
if err := r.cmd.Wait(); err != nil {
r.logger.Warnf("rclone cat process exited with error: %v", err)
}
diff --git a/storage/s3/s3.go b/storage/s3/s3.go
index d893c77..7c2e2be 100644
--- a/storage/s3/s3.go
+++ b/storage/s3/s3.go
@@ -8,11 +8,11 @@ import (
"strings"
"github.com/charmbracelet/log"
+ "github.com/krau/SaveAny-Bot/common/utils/fsutil"
storconfig "github.com/krau/SaveAny-Bot/config/storage"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
"github.com/krau/SaveAny-Bot/pkg/s3"
- "github.com/rs/xid"
)
type S3 struct {
@@ -65,21 +65,13 @@ func (m *S3) JoinStoragePath(p string) string {
func (m *S3) Save(ctx context.Context, r io.Reader, storagePath string) error {
m.logger.Infof("Saving file from reader to %s", storagePath)
- storagePath = m.JoinStoragePath(storagePath)
- ext := path.Ext(storagePath)
- base := strings.TrimSuffix(storagePath, ext)
- candidate := storagePath
+ candidate := m.JoinStoragePath(storagePath)
if overwrite, _ := ctx.Value(ctxkey.OverwriteExisting).(bool); !overwrite {
// Unique filename
- for i := 1; m.existsKey(ctx, candidate); i++ {
- candidate = fmt.Sprintf("%s_%d%s", base, i, ext)
- if i > 10 {
- m.logger.Errorf("Too many attempts for unique filename: %s", storagePath)
- candidate = fmt.Sprintf("%s_%s%s", base, xid.New().String(), ext)
- break
- }
- }
+ candidate = fsutil.UniquePath(strings.TrimPrefix(m.config.BasePath, "/"), storagePath, func(c string) bool {
+ return m.existsKey(ctx, c)
+ }, 10)
}
// Determine content length
diff --git a/storage/storage.go b/storage/storage.go
index 20936a4..1714350 100644
--- a/storage/storage.go
+++ b/storage/storage.go
@@ -75,11 +75,18 @@ type StorageReadable interface {
OpenFile(ctx context.Context, filePath string) (io.ReadCloser, int64, error)
}
-var Storages = make(map[string]Storage)
-
var _ StorageProgressSaver = (*telegram.Telegram)(nil)
var _ StorageBatchProgressSaver = (*telegram.Telegram)(nil)
+var _ StorageListable = (*alist.Alist)(nil)
+var _ StorageReadable = (*alist.Alist)(nil)
+var _ StorageListable = (*local.Local)(nil)
+var _ StorageReadable = (*local.Local)(nil)
+var _ StorageListable = (*rclone.Rclone)(nil)
+var _ StorageReadable = (*rclone.Rclone)(nil)
+var _ StorageListable = (*webdav.Webdav)(nil)
+var _ StorageReadable = (*webdav.Webdav)(nil)
+
type StorageConstructor func() Storage
var storageConstructors = map[storenum.StorageType]StorageConstructor{
diff --git a/storage/telegram/telegram.go b/storage/telegram/telegram.go
index 439f2f4..0a40a85 100644
--- a/storage/telegram/telegram.go
+++ b/storage/telegram/telegram.go
@@ -83,6 +83,9 @@ func (t *Telegram) Name() string {
return t.config.Name
}
+// Exists always reports false: Telegram offers no reliable way to query
+// whether a file already exists in a chat, so conflict policies do not apply
+// to this backend.
func (t *Telegram) Exists(ctx context.Context, storagePath string) bool {
return false
}
@@ -447,7 +450,7 @@ func planMediaGroups(items []batchMediaItem) [][]batchMediaItem {
continue
}
end := i + 1
- for end < len(items) && end-i < 10 {
+ for end < len(items) && end-i < tglimit.MaxAlbumItems {
next := items[end]
if next.useSingleSave || !next.albumEligible || next.chatID != item.chatID || next.item.SourceGroupKey != item.item.SourceGroupKey {
break
@@ -623,16 +626,16 @@ func (t *Telegram) splitUpload(
sender := ctx.Sender
- if len(multiMedia) <= 10 {
+ if len(multiMedia) <= tglimit.MaxAlbumItems {
_, err = sender.WithUploader(upler).
To(peer).
Album(ctx, multiMedia[0], multiMedia[1:]...)
return err
}
- // more than 10 parts, send in batches, each batch up to 10 parts
- for i := 0; i < len(multiMedia); i += 10 {
- end := min(i+10, len(multiMedia))
+ // more than MaxAlbumItems parts, send in batches, each batch up to MaxAlbumItems parts
+ for i := 0; i < len(multiMedia); i += tglimit.MaxAlbumItems {
+ end := min(i+tglimit.MaxAlbumItems, len(multiMedia))
batch := multiMedia[i:end]
_, err = sender.WithUploader(upler).
To(peer).
diff --git a/storage/telegram/video_split.go b/storage/telegram/video_split.go
index a4abc3d..afb5c55 100644
--- a/storage/telegram/video_split.go
+++ b/storage/telegram/video_split.go
@@ -19,13 +19,14 @@ import (
"github.com/rs/xid"
"github.com/krau/SaveAny-Bot/config"
+ "github.com/krau/SaveAny-Bot/pkg/consts/tglimit"
)
const (
videoPartTargetRatio = 0.95
videoSplitAttempts = 4
minSegmentDuration = 1.0
- maxLosslessVideoParts = 10
+ maxLosslessVideoParts = tglimit.MaxAlbumItems
)
type losslessVideoPart struct {
diff --git a/storage/webdav/errs.go b/storage/webdav/errs.go
deleted file mode 100644
index 9f9cd5b..0000000
--- a/storage/webdav/errs.go
+++ /dev/null
@@ -1,9 +0,0 @@
-package webdav
-
-import "errors"
-
-var (
- ErrFailedToCreateDirectory = errors.New("webdav: failed to create directory")
- ErrFailedToWriteFile = errors.New("webdav: failed to write file")
- ErrFailedToCheckFileExists = errors.New("webdav: failed to check if file exists")
-)
diff --git a/storage/webdav/webdav.go b/storage/webdav/webdav.go
index d869a5d..5fa61f8 100644
--- a/storage/webdav/webdav.go
+++ b/storage/webdav/webdav.go
@@ -11,11 +11,11 @@ import (
"time"
"github.com/charmbracelet/log"
+ "github.com/krau/SaveAny-Bot/common/utils/fsutil"
config "github.com/krau/SaveAny-Bot/config/storage"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
storenum "github.com/krau/SaveAny-Bot/pkg/enums/storage"
"github.com/krau/SaveAny-Bot/pkg/storagetypes"
- "github.com/rs/xid"
)
type Webdav struct {
@@ -54,28 +54,18 @@ func (w *Webdav) JoinStoragePath(p string) string {
func (w *Webdav) Save(ctx context.Context, r io.Reader, storagePath string) error {
w.logger.Infof("Saving file to %s", storagePath)
- storagePath = w.JoinStoragePath(storagePath)
- ext := path.Ext(storagePath)
- base := strings.TrimSuffix(storagePath, ext)
- candidate := storagePath
+ candidate := w.JoinStoragePath(storagePath)
if overwrite, _ := ctx.Value(ctxkey.OverwriteExisting).(bool); !overwrite {
- for i := 1; w.existsPath(ctx, candidate); i++ {
- candidate = fmt.Sprintf("%s_%d%s", base, i, ext)
- if i > 1000 {
- w.logger.Errorf("Too many attempts to find a unique filename for %s", storagePath)
- candidate = fmt.Sprintf("%s_%s%s", base, xid.New().String(), ext)
- break
- }
- }
+ candidate = fsutil.UniquePath(w.config.BasePath, storagePath, func(c string) bool {
+ return w.existsPath(ctx, c)
+ }, 1000)
}
if err := w.client.MkDir(ctx, path.Dir(candidate)); err != nil {
- w.logger.Errorf("Failed to create directory %s: %v", path.Dir(candidate), err)
- return ErrFailedToCreateDirectory
+ return fmt.Errorf("failed to create directory: %w", err)
}
if err := w.client.WriteFile(ctx, candidate, r); err != nil {
- w.logger.Errorf("Failed to write file %s: %v", candidate, err)
- return ErrFailedToWriteFile
+ return fmt.Errorf("failed to write file: %w", err)
}
return nil
}
@@ -136,9 +126,6 @@ func (w *Webdav) ListFiles(ctx context.Context, dirPath string) ([]storagetypes.
isDir := resp.Propstat.Prop.ResourceType.IsCollection()
- filePath := strings.TrimPrefix(decodedHref, path.Join("/", strings.Trim(path.Dir(fullPath), "/")))
- filePath = strings.TrimPrefix(filePath, "/")
-
fileInfo := storagetypes.FileInfo{
Name: name,
Path: path.Join(dirPath, name),