Compare commits

..

1 Commits

Author SHA1 Message Date
krau
df3c568bb8 feat: implement task event system for progress tracking and reporting 2026-06-25 17:52:38 +08:00
151 changed files with 1638 additions and 8547 deletions

View File

@@ -43,7 +43,7 @@ jobs:
goarch: arm64
steps:
- name: Checkout
uses: actions/checkout@v6
uses: actions/checkout@v4
- 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=${{ github.event.repository.updated_at }}"
-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.GitCommit=${{ github.sha }}"
binary_name: saveany-bot
env:

1
.gitignore vendored
View File

@@ -1,6 +1,5 @@
config.toml
logs/
/cache/
tmp/
data/
downloads/

376
AGENTS.md
View File

@@ -1,115 +1,301 @@
# Repository Guidelines
# SaveAny-Bot Agent Guidelines
This document provides essential information for AI coding agents working on the SaveAny-Bot project.
## Project Overview
SaveAny-Bot is a Telegram bot written in Go that saves files and messages from Telegram and websites to multiple storage backends (Local, S3, MinIO, WebDAV, AList, Rclone, Telegram). It supports single-file saves, batch/album saves, streaming, multi-user access, storage rules, cross-storage transfers, yt-dlp and Aria2 downloads, and a Goja-based JavaScript parser plugin system with optional Playwright browser automation.
SaveAny-Bot is a Telegram bot written in Go that saves files/messages from Telegram and various websites to multiple storage backends (local, S3, MinIO, WebDAV, AList, Telegram). It features a plugin system for parsing web content and extensible storage backends.
**Tech stack**: Go 1.25, gotgproto + gotd/td v0.149.0 (MTProto), Cobra (CLI), Viper (config), GORM + SQLite, Goja (JS runtime), Playwright, charmbracelet/log. License: AGPL-3.0.
**Tech Stack**: Go 1.24.2, gotd/td (Telegram MTProto), Cobra (CLI), Viper (config), GORM (ORM), SQLite, Goja (JS runtime), Playwright (browser automation)
**Note on gotd versions**: `gotd/td` must stay on `v0.149.0` — v0.150+ breaks `gotgproto` (v1.0.0-beta22) compilation (`AsInputDocumentFileLocation` signature, `gotd/log` Logger interface). Do not bump beyond v0.149.
## Architecture & Data Flow
```
Telegram update → client/bot/handlers → core.AddTask(Executable) → pkg/queue (serial workers)
→ core/tasks/* (download via common/tdler) → storage.Storage → progress feedback
```
- **Startup sequence** (`cmd/run.go::initAll`, keep this order): Config → Cache → i18n → Database → Storage → Parser plugins → Userbot → API → Bot. `bot.Init` returns the exit channel; a `SAVEANTBOT-RESTART` error restarts the process (external supervisor).
- **Task pipeline**: handlers build a task (via `core.AddTask`), the queue executes `Executable{Type, Title, TaskID, Execute(ctx)}` with `config.C().Workers` workers. Lifecycle hooks (`TaskBeforeStart/Success/Fail/Cancel`) run around `Execute`. Cancellation = canceling the task's context; tasks must check `ctx.Err()`.
- **Dual progress channels**: (1) `pkg/taskevent` context bus (consumed by `api/` for HTTP/Webhook consumers — `taskevent.WithSink` must be injected for API tasks), (2) Telegram message edits via `ProgressTracker` + `tgutil.ExtFromContext(ctx)` (bot tasks). Upload progress currently reaches the Telegram channel only.
- **Capability interfaces + type assertion fallback** is the core extensibility pattern: `StorageBatchSaver`, `StorageProgressSaver` (upload progress), `StorageListable`, `StorageReadable` are optional; consumers assert and fall back (e.g. wrap reader with `ioutil.NewProgressReader`). New backends only need to implement the interface and register.
- **Config layering**: CLI flag > env `SAVEANY_*` (dots → underscores, e.g. `SAVEANY_TELEGRAM_TOKEN`) > TOML file (path or http(s) URL). `config.C()` returns a **copy** — never mutate it.
- **Storage registration (3 places)**: `pkg/enums/storage` ENUM comment (go-enum), `config/storage/factory.go::storageFactories` (config struct with `Validate()`), `storage/storage.go::storageConstructors` (implementation).
## Key Directories
| Path | Purpose |
|---|---|
| `cmd/` | CLI: `run` (main bot), `upload`, `watch` (standalone subcommands that do NOT run initAll), `geni18n` (i18n key generator) |
| `core/` | `Executable` interface, queue worker loop, hooks; `core/tasks/{tfile,batchtfile,directlinks,parsed,telegraph,transfer,ytdlp,aria2dl}` |
| `client/bot/` | gotgproto client, `handlers/` (all commands + message/callback handlers), `middleware/`, `client/user/` (userbot) |
| `storage/` | 8 backends + `storage.go` (interfaces/registry) + `load.go` (per-user storage resolution) |
| `parsers/` | `parsers.go` (registry), `js/` (Goja plugins, ghttp/playwright injection, build-tagged), `parsers/` (native: twitter, kemono) |
| `config/` | Viper setup, defaults, `storage/` per-type config structs |
| `database/` | GORM models (User/Dir/Rule/WatchChat), AutoMigrate, `syncUsers` |
| `pkg/` | `queue`, `taskevent`, `tcbdata` (callback data), `rule`, `enums/{tasktype,storage,ctxkey,fnamest}`, `storagetypes`, `tfile`, `parser` |
| `common/` | `tdler` (unified downloader), `utils/{tgutil,dlutil,ioutil,fsutil,strutil,tphutil,netutil}`, `i18n` (embedded locales), `cache` (ristretto) |
| `api/` | HTTP API + webhook (task factory with sink injection) |
| `docs/` | Hugo site (hugo-book theme, zh+en mirrored), separate go.mod |
| `plugins/` | JS parser examples + `README.md` (plugin author contract) |
## Development Commands
## Build & Test Commands
### Build
```bash
# Build (standard; CGO_ENABLED=0 for static)
CGO_ENABLED=0 go build -trimpath -o saveany-bot .
# Standard build
go build -o saveany-bot .
# Run directly
go run ./cmd
# Test — known failures: storage/telegram TestCreateSplitZip/TestExtractThumbFrame/TestGetVideoMetadata
# (need gitignored fixtures tests/testfile.dat, tests/testvideo; ffmpeg/ffprobe)
go test ./...
go test -race ./core/tasks/... ./storage/... ./pkg/queue/... ./common/...
go test -run TestQueueBasic ./pkg/queue
# Codegen — run after editing locale YAML or enum comments
go generate ./... # geni18n (i18nk keys) + go-enum (pkg/enums/*)
# go-enum is NOT in go.mod; install externally. geni18n runs via go run.
# Verify
go vet ./...
go fmt ./...
# Docker build (multi-stage, Alpine-based)
docker build -t saveany-bot .
docker compose up -d
```
**Build variants** (Dockerfile.default/micro/pico): `-tags=no_jsparser,no_playwright,no_minio,no_bubbletea,sqlite_glebarez` — each has a `*_stub.go`/`*_glebarez.go` pairing; keep stubs in sync.
### Test
```bash
# Run all tests
go test ./...
Docker: `docker build -t saveany-bot .`, `docker compose up -d` (host network, mounts `./data ./config.toml ./downloads ./cache`). CI (`.github/workflows/`) runs **no tests/lint** — only tag-triggered release/docker builds and docs deployment; run `go test ./...` manually before pushing.
# Run tests in specific package
go test ./pkg/queue
go test ./storage/telegram
## Code Conventions & Common Patterns
# Run tests with verbose output
go test -v ./...
- **Imports**: stdlib → third-party → project-internal, blank-line separated. Aliases for clarity (`storconfig`, `storenum`).
- **Naming**: PascalCase exported, camelCase unexported, files `snake_case.go`; **not** ALL_CAPS constants.
- **Errors**: always wrap with `fmt.Errorf("context: %w", err)`; check with `errors.Is/As`; never ignore.
- **Logging**: `log.FromContext(ctx)` with prefixes (`logger.WithPrefix("component")`); never global logger when ctx is available.
- **Context values** (read from the passed ctx, never globals): `log.FromContext`, `tgutil.ExtFromContext` (Telegram ext — **required for message edits; if nil, edits are silently dropped**), `storage.FromContext`, `storagetypes.WithSourceCaption`, `ctxkey.ContentLength` / `ctxkey.OverwriteExisting`.
- **Progress rendering** (#228 convention): i18n templates declare styles with Telegram HTML (`<b>/<code>/<blockquote>/<i>`); dynamic data MUST go through `i18n.T(key, tgutil.EscapeHTMLTemplateData(data))` before `tgutil.RenderHTML`. Never interpolate user data raw, never render-then-substring-search.
- **Progress tracking**: each task package defines its own small `ProgressTracker` interface; optional `UploadProgressTracker` is probed via type assertion (skip if absent). Serialize state + message edits with a mutex; throttle edits (≥1s); aggregate per-item progress monotonically.
- **i18n**: only edit `common/i18n/locale/{zh-Hans,en}.yaml``go generate ./...` → use `i18nk.<Key>` constants. No raw strings in user-facing messages. zh-Hans and en must stay in sync.
- **Registration points** (never forget): new bot command → `client/bot/handlers/register.go::CommandHandlers` (auto-publishes /help menu); new task type → `pkg/enums/tasktype` + `core/tasks/<name>/` + `api/factory.go::CreateTask`; new storage → 3 places above + `docs/content/{en,zh}/deployment/configuration/storages.md`; new enum value → ENUM comment + `go generate`.
- **Concurrency**: `errgroup.WithContext` + `SetLimit(config.C().Workers)`, `atomic.Int64` counters, `sync.Once` for single-shot events, mutex around render state. No lock-in-callback (callbacks fire after unlock).
- **Cancellation**: queue tasks carry a `WithCancel`-derived ctx; check `ctx.Err()` in loops; classify with `errors.Is(err, context.Canceled)`.
- **JS plugins**: `registerParser({metadata, canHandle, parse})`, `version >= 1.0.0`; per-plugin goja VM is single-goroutine (reqCh buffer 10). Changing `pkg/parser.Item/Resource` JSON fields requires updating `plugins/README.md` and example plugins.
- **Message edits**: `ext.EditMessage(chatID, &tg.MessagesEditMessageRequest{...})`; cancel buttons via `tgutil.BuildCancelButton(taskID)`; callback payloads via `pkg/tcbdata` + `common/cache`.
# Run a single test
go test -run TestQueueBasic ./pkg/queue
## Important Files
# Run with coverage
go test -cover ./...
```
- `main.go``//go:generate` for i18n keys
- `cmd/run.go` — startup sequence `Run/initAll/cleanCache` (cache cleanup on exit, `NoCleanCache` opt-out)
- `core/core.go` — worker loop, hooks, AddTask/CancelTask
- `pkg/queue/queue.go` — generic serial queue (cond/list; duplicate TaskID rejected)
- `storage/storage.go` — interfaces + registry + compile-time capability assertions
- `config/viper.go`, `config.example.toml` — config schema (authoritative field docs)
- `database/db.go` — GORM init, `GetDialect` (build-tag selectable SQLite driver)
- `client/bot/handlers/register.go` — handler dispatch order and CommandHandlers
- `common/tdler/dler.go` — unified download entry
- `core/tasks/batchtfile/item_progress.go` — per-item phase state machine (Downloading/Transferring/Uploading/Retrying/Confirming, FailureStage)
- `parsers/js/plugin.go` — Goja plugin runtime
- `.github/workflows/` — release/docker/docs (no test gate)
### Lint & Format
```bash
# Format code (standard Go formatting)
go fmt ./...
## Runtime/Tooling Preferences
# Vet code for common issues
go vet ./...
- **Go 1.25+**: `t.Context()`, `sync.WaitGroup.Go`, `for range n` are available.
- **Runtime binaries**: ffmpeg/ffprobe (media processing/video split), yt-dlp (ytdlp tasks), aria2 optional; Playwright browsers install on demand to `./playwright` (`playwright.Install(chromium, ...)` at first `pw.get()`); Docker images: default has ffmpeg+yt-dlp, micro only curl, pico is scratch static.
- **No Makefile, no golangci.yml, no test/lint CI** — verification is manual (`go vet`, `go test`).
- **go-enum** required externally for enum generation; **geni18n** is in-repo.
- **Docs**: Hugo site in `docs/` (separate go.mod, hugo-book); edit `docs/content/{zh,en}/` — keep both languages mirrored. `docs/public/` is gitignored build output.
- **gitignored fixtures**: `storage/telegram/tests/` (missing — 3 tests fail locally), `data/`, `config.toml`, `playwright/`, `testplugins/`.
# Generate code (i18n keys)
go generate ./...
```
## Testing & QA
### Other Commands
```bash
# Update dependencies
go mod tidy
- Pure stdlib `testing` (no testify); table-driven (`[]struct{name...}` + `t.Run`) with `t.Fatalf` got/want assertions. Mock via hand-written interface impls or package-variable replacement (`runMediaTool` in `video_split_test.go`, restored with `t.Cleanup`); in-process services for HTTP (`httptest`), S3 (`gofakes3+s3mem`), WebDAV (`x/net/webdav`).
- **Locale-dependent tests**: pin with `i18n.Init("zh-Hans")` + `t.Cleanup(...)`.
- **Progress/HTML tests**: assert rendered text with `strings.Contains` AND entity counts (`tg.MessageEntityBold/Code/Blockquote/Italic`) — verify style injection stays escaped (`<b>A&B</b>` input must render as literal text).
- **Known failures**: `storage/telegram` `TestCreateSplitZip`, `TestExtractThumbFrame`, `TestGetVideoMetadata` need gitignored fixtures + real ffmpeg — skip with `-skip 'Test(CreateSplitZip|ExtractThumbFrame|GetVideoMetadata)$'`; `api/handlers_test.go` has one `t.Skip` (needs initialized core).
- **Coverage expectations**: pure logic gets table tests (parsers, URL/path utils, progress throttling, grouping); regressions get bug-scenario-named tests (`progress_regression_test.go`). Network/Telegram/Playwright must never be touched by tests.
- When a permanent feature/API change ships: update `config.example.toml` if config, `docs/` if user-facing, `plugins/README.md` if plugin contract, and i18n YAML + `go generate` for new strings.
# View documentation
cd docs && hugo server -D
```
## Code Style Guidelines
### Imports
- Standard library first, then third-party, then project-internal
- Group imports with blank lines between groups
- Use explicit import aliases for clarity when needed (e.g., `storconfig`, `storenum`)
```go
import (
"context"
"fmt"
"github.com/charmbracelet/log"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/pkg/enums/storage"
)
```
### Formatting
- Line length: reasonable (no hard limit, but be sensible)
- Organize code with blank lines between logical sections
- Follow standard Go conventions for braces, spacing, etc.
### Types & Interfaces
- Use clear, descriptive type names (PascalCase for exported, camelCase for unexported)
- Define interfaces where abstraction is needed (e.g., `Executable`, `StorageConfig`)
- Embed context in method signatures, not structs: `func (s *Service) Do(ctx context.Context) error`
- Prefer composition over inheritance
```go
// Interfaces define behavior
type Executable interface {
Type() tasktype.TaskType
Title() string
TaskID() string
Execute(ctx context.Context) error
}
// Structs compose behavior
type Local struct {
config config.LocalStorageConfig
logger *log.Logger
}
```
### Naming Conventions
- **Packages**: lowercase, single word when possible (avoid underscores)
- **Files**: lowercase with underscores for multiword (e.g., `auth_terminal.go`, `progress_reader.go`)
- **Variables**: camelCase for unexported, PascalCase for exported
- **Constants**: PascalCase for exported, camelCase for unexported (not ALL_CAPS)
- **Functions/Methods**: PascalCase for exported, camelCase for unexported
- **Test files**: `*_test.go` pattern
### Error Handling
- Always handle errors explicitly; never ignore them
- Wrap errors with context using `fmt.Errorf("context: %w", err)`
- Use `errors.Is()` and `errors.As()` for error checking
- Log errors with appropriate level (Error, Warn, Info)
- Return errors from functions rather than panicking (except for truly unrecoverable situations)
```go
// Good error handling
if err := db.Save(user).Error; err != nil {
return fmt.Errorf("failed to save user %d: %w", user.ChatID, err)
}
// Check specific errors
if errors.Is(err, context.Canceled) {
logger.Info("Operation was canceled")
return nil
}
```
### Logging
- Use `github.com/charmbracelet/log` package
- Get logger from context: `log.FromContext(ctx)`
- Create prefixed loggers for components: `logger.WithPrefix("component")`
- Use appropriate levels: Debug, Info, Warn, Error
- Include context in log messages (e.g., task IDs, file names)
```go
logger := log.FromContext(ctx)
logger.Infof("Processing task: %s", task.ID)
logger.Errorf("Failed to save file %s: %v", filename, err)
```
### Concurrency
- Use channels for communication between goroutines
- Protect shared state with `sync.Mutex` or `sync.RWMutex`
- Use `sync.WaitGroup` for coordinating goroutine completion
- Always pass `context.Context` for cancellation support
- Use `context.WithCancel/WithTimeout` for managing goroutine lifetimes
```go
// Example from queue implementation
func (tq *TaskQueue[T]) Add(task *Task[T]) error {
tq.mu.Lock()
defer tq.mu.Unlock()
// ... critical section
tq.cond.Signal()
return nil
}
```
### Comments
- Document exported types, functions, and packages with doc comments
- Start doc comments with the name being documented
- Use `//` for single-line comments
- Explain *why*, not *what* (code should be self-explanatory for "what")
- Add `[NOTE]`, `[WARN]`, `[IMPORTANT]` tags for important clarifications
```go
// GetUserByChatID retrieves a user by their Telegram chat ID.
// Returns an error if the user is not found.
func GetUserByChatID(ctx context.Context, chatID int64) (*User, error) {
```
## Architecture & Conventions
### Application Structure
- **Entry point**: `main.go``cmd.Execute(ctx)`
- **CLI root**: `cmd/root.go` (Cobra), implementation in `cmd/run.go`
- **Startup sequence**: Config → Cache → i18n → Database → Storage → Parsers → Userbot → Bot → Queue
- Follow this order when adding new initialization steps in `cmd/run.go::initAll`
### Configuration (Viper)
- Config defined in `config/viper.go::Config`
- Read from `config.toml` (see `config.example.toml`)
- Environment variables: `SAVEANY_*` prefix (e.g., `SAVEANY_TELEGRAM_TOKEN`)
- Access via `config.C()` (returns a copy, don't modify the return value)
- Storage configs validated via `config/storage/factory.go::LoadStorageConfigs`
### Telegram Client
- **Bot client**: `client/bot/bot.go::Init` (uses gotgproto)
- **Handlers**: Centralized in `client/bot/handlers/` directory
- **Registration**: All handlers registered in `handlers.Register`
- **Commands**: Add to `CommandHandlers` slice for automatic `/help` and bot command list updates
- **Middleware**: Common middleware in `client/middleware/` (floodwait, retry, etc.)
### Tasks & Queue
- **Task interface**: `core/core.go::Executable` (Type, Title, TaskID, Execute methods)
- **Queue**: `pkg/queue.TaskQueue[Executable]` (generic, thread-safe)
- **Workers**: Count from `config.C().Workers`
- **Task types**: Implementations in `core/tasks/**` (tfile, parsed, telegraph, directlinks, batchtfile)
- **Lifecycle hooks**: `TaskBeforeStart`, `TaskSuccess`, `TaskFail`, `TaskCancel` (defined in config)
- **Adding tasks**: Use `core.AddTask(ctx, task)`
### Database (GORM + SQLite)
- **Init**: `database.Init` using `config.C().DB.Path`
- **Models**: User, Dir, Rule, WatchChat (in `database/*.go`)
- **Migrations**: Automatic via `db.AutoMigrate`
- **User sync**: `database.syncUsers` syncs DB with `config.C().Users` (don't manually create/delete users)
- **Context**: Always use `db.WithContext(ctx)` for operations
### Storage Backends
- **Interface**: Defined in `config/storage/types.go` and `storage/`
- **Implementations**: local, alist, s3/minio, webdav, telegram (each in subdirectory)
- **Adding new storage**:
1. Add enum to `pkg/enums/storage`
2. Create config struct in `config/storage/` with `Validate()` method
3. Implement storage in `storage/<name>/`
4. Register in `storageFactories` mapping
5. Update `config.example.toml` with example
### Parser Plugins (JavaScript)
- **Runtime**: Goja (JS runtime) + Playwright (browser automation)
- **Plugin API**: `registerParser({ metadata, canHandle, parse })` in JS
- **Integration**: Defined in `parsers/` directory
- **Documentation**: See `plugins/README.md`
- Plugin `parse` returns `Item`/`Resource` which becomes download/transfer task
### Internationalization (i18n)
- **Usage**: `i18n.T(i18nk.SomeKey, map[string]any{"Name": value})`
- **Locale files**: `common/i18n/locale/*.yaml`
- **Key generation**: Run `go generate ./...` to generate `common/i18n/i18nk/keys.go`
- **Adding new strings**: Add to YAML → run `go generate` → use in code
- All user-facing strings should be internationalized
### Context Usage
- Always pass `context.Context` as first parameter
- Use `log.FromContext(ctx)` to get contextual logger
- Respect context cancellation in long-running operations
- Store request-scoped data in context (e.g., `ctxkey.ContentLength`)
## Special Rules from .github/copilot-instructions.md
1. **Never modify `config.C()` return values** - it returns a copy. Modify config in `config.Init` or via Viper.
2. **Handlers must update `CommandHandlers` slice** - ensures `/help` and bot commands stay in sync.
3. **Task execution must preserve hooks** - don't remove `TaskBeforeStart`, `TaskSuccess`, `TaskFail`, `TaskCancel` hook calls.
4. **User sync is automatic** - don't manually create/delete users in DB; use config-based sync.
5. **Prefer context logger** - use `log.FromContext(ctx)` over global logger when context is available.
6. **Storage factory pattern** - new storage types must register in `storageFactories` mapping.
7. **Plugin API compatibility** - changes to `Item`/`Resource` structures require updating `plugins/README.md`.
## Common Patterns
### Adding a New Command
1. Create handler function in `client/bot/handlers/<name>.go`
2. Add to `CommandHandlers` slice in `register.go`
3. Add i18n key to `common/i18n/locale/*.yaml`
4. Run `go generate ./...`
5. Test with Telegram bot
### Adding a New Task Type
1. Create struct implementing `core.Executable` in `core/tasks/<type>/`
2. Implement `Type()`, `Title()`, `TaskID()`, `Execute(ctx)` methods
3. Add task type enum to `pkg/enums/tasktype`
4. Use `core.AddTask(ctx, task)` to enqueue
### Adding a New Storage Backend
1. Define config struct in `config/storage/<name>.go` with `Validate()` method
2. Implement storage interface in `storage/<name>/<name>.go`
3. Add storage type enum to `pkg/enums/storage`
4. Register factory in `config/storage/factory.go::storageFactories`
5. Update `config.example.toml` with configuration example
## File References
When referencing code locations, use `path/to/file.go:line` format (e.g., `core/core.go:23` for the worker function).
## Testing Guidelines
- Write tests for new functionality (place in `*_test.go` files)
- Test files should be in same package as code being tested
- Use table-driven tests for multiple test cases
- Mock external dependencies (databases, network calls)
- Aim for meaningful tests, not just coverage numbers
## Notes
- Binary size matters: use `CGO_ENABLED=0` for static binaries
- FFmpeg is included in Docker images for media processing
- Build process supports cross-compilation (amd64/arm64, Linux/macOS/Windows)
- Documentation site uses Hugo; edit files in `docs/` directory
- Session data stored in SQLite; delete `data/session.db` if changing bot token

View File

@@ -7,8 +7,6 @@ ARG BuildTime="Unknown"
WORKDIR /app
RUN apk add --no-cache ca-certificates
COPY go.mod go.sum ./
RUN --mount=type=cache,target=/go/pkg/mod \
go mod download
@@ -33,9 +31,5 @@ FROM scratch
WORKDIR /app
COPY --from=builder /app/saveany-bot .
COPY --from=builder /etc/ssl/certs/ca-certificates.crt /etc/ssl/certs/ca-certificates.crt
ENV SSL_CERT_FILE=/etc/ssl/certs/ca-certificates.crt
ENV SSL_CERT_DIR=/etc/ssl/certs
ENTRYPOINT ["/app/saveany-bot"]

View File

@@ -1,6 +1,7 @@
package api
import (
"context"
"crypto/subtle"
"net/http"
"strings"
@@ -8,6 +9,9 @@ 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 {
@@ -36,7 +40,9 @@ func AuthMiddleware() func(http.Handler) http.Handler {
return
}
next.ServeHTTP(w, r)
// 将 token 添加到 context
ctx := context.WithValue(r.Context(), tokenContextKey{}, token)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}

View File

@@ -39,7 +39,7 @@ func NewTaskFactory(ctx context.Context) *TaskFactory {
// CreateTask 创建任务
func (f *TaskFactory) CreateTask(req *CreateTaskRequest) (*CreateTaskResponse, error) {
// 验证存储
stor, ok := storage.GetStorage(req.Storage)
stor, ok := storage.Storages[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.GetStorage(params.SourceStorage)
sourceStor, ok := storage.Storages[params.SourceStorage]
if !ok {
return nil, fmt.Errorf("source storage not found: %s", params.SourceStorage)
}
targetStor, ok := storage.GetStorage(params.TargetStorage)
targetStor, ok := storage.Storages[params.TargetStorage]
if !ok {
return nil, fmt.Errorf("target storage not found: %s", params.TargetStorage)
}

View File

@@ -135,9 +135,8 @@ func (h *Handlers) ListStoragesHandler(w http.ResponseWriter, r *http.Request) {
return
}
all := storage.AllStorages()
storages := make([]StorageInfo, 0, len(all))
for name, stor := range all {
storages := make([]StorageInfo, 0, len(storage.Storages))
for name, stor := range storage.Storages {
storages = append(storages, StorageInfo{
Name: name,
Type: string(stor.Type()),

View File

@@ -1,7 +1,6 @@
package api
import (
"context"
"sync"
"time"
@@ -12,23 +11,23 @@ import (
// guarded by mu. It implements taskevent.Sink so the task layer can update it
// without knowing about the API.
type TaskProgressInfo struct {
mu sync.Mutex
TaskID string
Type string
Status TaskStatus
Title string
TotalBytes int64
DownloadedBytes int64
TotalFiles int
DownloadedFiles int
Storage string
Path string
Error string
CreatedAt time.Time
UpdatedAt time.Time
StartedAt time.Time
Webhook string
webhookNotified bool
mu sync.Mutex
TaskID string
Type string
Status TaskStatus
Title string
TotalBytes int64
DownloadedBytes int64
TotalFiles int
DownloadedFiles int
Storage string
Path string
Error string
CreatedAt time.Time
UpdatedAt time.Time
StartedAt time.Time
Webhook string
webhookNotified bool
}
// progressStore holds all API tasks. Entries are removed a fixed duration after
@@ -196,6 +195,22 @@ func (t *TaskProgressInfo) Emit(e taskevent.Event) {
if notify {
payload := CreateWebhookPayload(t.TaskID, t.Type, t.Status, t.Storage, t.Path, e.Err)
SendWebhook(context.Background(), payload)
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) {}

View File

@@ -3,7 +3,6 @@ package api
import (
"context"
"fmt"
"net"
"net/http"
"time"
@@ -91,15 +90,9 @@ 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.Serve(ln); err != nil && err != http.ErrServerClosed {
if err := s.httpServer.ListenAndServe(); err != nil && err != http.ErrServerClosed {
logger.Errorf("API server error: %v", err)
}
}()

View File

@@ -60,7 +60,6 @@ func resolveChatID(_ context.Context, idOrUsername string) (int64, error) {
}
// ParseMessageLink 解析 Telegram 消息链接
// 支持的域名: t.me, telegram.me
// 支持格式:
// - https://t.me/username/123
// - https://t.me/c/123456789/123
@@ -269,15 +268,5 @@ func ExtractFilesFromLinks(ctx context.Context, links []string) ([]tfile.TGFileM
// isValidMessageLink 检查是否是有效的 Telegram 消息链接
func isValidMessageLink(link string) bool {
for _, prefix := range []string{
"https://t.me/",
"http://t.me/",
"https://telegram.me/",
"http://telegram.me/",
} {
if strings.HasPrefix(link, prefix) {
return true
}
}
return false
return strings.HasPrefix(link, "https://t.me/") || strings.HasPrefix(link, "http://t.me/")
}

View File

@@ -37,9 +37,6 @@ 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 {
@@ -47,15 +44,10 @@ func SendWebhook(ctx context.Context, payload *WebhookPayload) {
return
}
// 重试 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))
// 重试 3 次
for i := range 3 {
req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, webhookURL, bytes.NewBuffer(payloadBytes))
if err != nil {
cancel()
logger.Errorf("Failed to create webhook request: %v", err)
return
}
@@ -64,13 +56,9 @@ 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/%d): %v", i+1, maxAttempts, err)
if i < maxAttempts-1 {
time.Sleep(backoff)
}
backoff *= 4
logger.Warnf("Webhook request failed (attempt %d/3): %v", i+1, err)
time.Sleep(time.Second * time.Duration(i+1))
continue
}
resp.Body.Close()
@@ -80,14 +68,11 @@ func SendWebhook(ctx context.Context, payload *WebhookPayload) {
return
}
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.Warnf("Webhook returned non-2xx status (attempt %d/3): %d", i+1, resp.StatusCode)
time.Sleep(time.Second * time.Duration(i+1))
}
logger.Errorf("Failed to send webhook after %d attempts", maxAttempts)
logger.Errorf("Failed to send webhook after 3 attempts")
}()
}

View File

@@ -70,6 +70,9 @@ 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)})

View File

@@ -23,11 +23,7 @@ import (
)
func handleAddCallback(ctx *ext.Context, update *ext.Update) error {
dataParts := strings.Split(string(update.CallbackQuery.Data), " ")
if len(dataParts) < 2 {
return fmt.Errorf("invalid callback data: %q", update.CallbackQuery.Data)
}
dataid := dataParts[1]
dataid := strings.Split(string(update.CallbackQuery.Data), " ")[1]
data, err := shortcut.GetCallbackDataWithAnswer[tcbdata.Add](ctx, update, dataid)
if err != nil {
return err

View File

@@ -1,7 +1,6 @@
package handlers
import (
"fmt"
"strings"
"github.com/celestix/gotgproto/dispatcher"
@@ -15,11 +14,7 @@ import (
)
func handleCancelCallback(ctx *ext.Context, update *ext.Update) error {
dataParts := strings.Split(string(update.CallbackQuery.Data), " ")
if len(dataParts) < 2 {
return fmt.Errorf("invalid callback data: %q", update.CallbackQuery.Data)
}
taskid := dataParts[1]
taskid := strings.Split(string(update.CallbackQuery.Data), " ")[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{

View File

@@ -10,7 +10,6 @@ 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"
)
@@ -43,9 +42,7 @@ func handleDirCmd(ctx *ext.Context, update *ext.Update) error {
return dispatcher.EndGroups
}
if _, err := storage.GetStorageByUserIDAndName(ctx, user.ChatID, args[2]); err != 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)
ctx.Reply(update, ext.ReplyTextString(err.Error()), nil)
return dispatcher.EndGroups
}

View File

@@ -13,24 +13,16 @@ import (
"github.com/krau/SaveAny-Bot/client/bot/handlers/utils/shortcut"
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/database"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/pkg/tcbdata"
"github.com/krau/SaveAny-Bot/pkg/tfile"
"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[mediaGroupKey][]tfile.TGFileMessage
timers map[mediaGroupKey]*time.Timer
groups map[int64][]tfile.TGFileMessage
timers map[int64]*time.Timer
mu sync.Mutex
timeout time.Duration
setupOnce sync.Once
@@ -47,8 +39,8 @@ func (m *MediaGroupHandler) SetupTimeout(timeoutSec int) {
var (
mediaGroupHandler = &MediaGroupHandler{
groups: make(map[mediaGroupKey][]tfile.TGFileMessage),
timers: make(map[mediaGroupKey]*time.Timer),
groups: make(map[int64][]tfile.TGFileMessage),
timers: make(map[int64]*time.Timer),
mu: sync.Mutex{},
}
)
@@ -74,37 +66,32 @@ func handleGroupMediaMessage(ctx *ext.Context, update *ext.Update, message *tg.M
}
mediaGroupHandler.mu.Lock()
defer mediaGroupHandler.mu.Unlock()
key := mediaGroupKey{
chatID: update.EffectiveChat().GetID(),
userID: userId,
groupID: groupID,
if mediaGroupHandler.groups[groupID] == nil {
mediaGroupHandler.groups[groupID] = make([]tfile.TGFileMessage, 0)
}
if mediaGroupHandler.groups[key] == nil {
mediaGroupHandler.groups[key] = make([]tfile.TGFileMessage, 0)
}
mediaGroupHandler.groups[key] = append(mediaGroupHandler.groups[key], file)
mediaGroupHandler.groups[groupID] = append(mediaGroupHandler.groups[groupID], file)
if timer, exists := mediaGroupHandler.timers[key]; exists {
if timer, exists := mediaGroupHandler.timers[groupID]; exists {
timer.Stop()
}
mediaGroupHandler.timers[key] = time.AfterFunc(mediaGroupHandler.timeout, func() {
processMediaGroup(ctx, update, key)
mediaGroupHandler.timers[groupID] = time.AfterFunc(mediaGroupHandler.timeout, func() {
processMediaGroup(ctx, update, groupID)
})
return dispatcher.EndGroups
}
func processMediaGroup(ctx *ext.Context, update *ext.Update, key mediaGroupKey) {
func processMediaGroup(ctx *ext.Context, update *ext.Update, groupID int64) {
logger := log.FromContext(ctx)
mediaGroupHandler.mu.Lock()
items := mediaGroupHandler.groups[key]
delete(mediaGroupHandler.groups, key)
delete(mediaGroupHandler.timers, key)
items := mediaGroupHandler.groups[groupID]
delete(mediaGroupHandler.groups, groupID)
delete(mediaGroupHandler.timers, groupID)
mediaGroupHandler.mu.Unlock()
if len(items) == 0 {
logger.Warn("No media items to process for group", "groupID", key.groupID)
logger.Warn("No media items to process for group", "groupID", groupID)
return
}
logger.Debugf("Processing media group %d with %d items", key.groupID, len(items))
logger.Debugf("Processing media group %d with %d items", groupID, len(items))
userId := update.GetUserChat().GetID()
msg, err := ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgMediaGroupInfoSavingFiles, nil)), nil)

View File

@@ -1,13 +1,10 @@
package handlers
import (
"errors"
"github.com/celestix/gotgproto/dispatcher"
"github.com/celestix/gotgproto/ext"
"github.com/duke-git/lancet/v2/slice"
"github.com/krau/SaveAny-Bot/client/bot/handlers/utils/dirutil"
"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"
@@ -15,42 +12,16 @@ import (
"github.com/krau/SaveAny-Bot/storage"
)
// responsibleUserID returns the sender's ID. Callback queries carry it
// natively; message updates resolve it through the entity map.
func responsibleUserID(u *ext.Update) int64 {
if u.CallbackQuery != nil {
return u.CallbackQuery.GetUserID()
}
return u.GetUserChat().GetID()
}
func checkPermission(ctx *ext.Context, update *ext.Update) error {
userID := responsibleUserID(update)
userID := update.GetUserChat().GetID()
if !slice.Contain(config.C().GetUsersID(), userID) {
if cbq := update.CallbackQuery; cbq != nil {
ctx.AnswerCallback(msgelem.AlertCallbackAnswer(cbq.GetQueryID(), i18n.T(i18nk.BotMsgCommonErrorNoPermission, nil)))
} else {
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgCommonErrorNoPermission, nil)), nil)
}
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgCommonErrorNoPermission, nil)), nil)
return dispatcher.EndGroups
}
return dispatcher.ContinueGroups
}
// withPermission wraps a callback handler with the same whitelist check used
// for message handlers (checkPermission). ContinueGroups is the dispatcher's
// success sentinel, not an error: only real failures and EndGroups stop the
// chain before the wrapped handler runs.
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 && !errors.Is(err, dispatcher.ContinueGroups) {
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()

View File

@@ -1,85 +0,0 @@
package handlers
import (
"os"
"path/filepath"
"testing"
"github.com/celestix/gotgproto/ext"
"github.com/celestix/gotgproto/types"
"github.com/gotd/td/tg"
"github.com/krau/SaveAny-Bot/config"
)
// Regression: callback queries usually arrive as updateShort without entity
// maps, so resolving the sender through the entity map yields ID 0 and every
// click was denied by the whitelist check. Callback updates must use the
// native UserID field.
func TestResponsibleUserID(t *testing.T) {
tests := []struct {
name string
update *ext.Update
want int64
}{
{
name: "callback query uses native user id",
update: &ext.Update{CallbackQuery: &tg.UpdateBotCallbackQuery{UserID: 42}},
want: 42,
},
{
name: "message resolves through entity map",
update: &ext.Update{
EffectiveMessage: &types.Message{Message: &tg.Message{PeerID: &tg.PeerUser{UserID: 7}}},
Entities: &tg.Entities{Users: map[int64]*tg.User{7: {ID: 7}}},
},
want: 7,
},
{
name: "callback query ignores entity map",
update: &ext.Update{
CallbackQuery: &tg.UpdateBotCallbackQuery{UserID: 9},
Entities: &tg.Entities{Users: map[int64]*tg.User{8: {ID: 8}}},
},
want: 9,
},
{
name: "unresolvable update yields zero",
update: &ext.Update{},
want: 0,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := responsibleUserID(tt.update); got != tt.want {
t.Fatalf("responsibleUserID() = %d, want %d", got, tt.want)
}
})
}
}
// Regression: withPermission must treat ContinueGroups (the dispatcher's
// success sentinel) as a pass and invoke the wrapped handler. v0.60.1 treated
// it as an error, so every permitted callback was swallowed before the real
// handler ran.
func TestWithPermissionInvokesHandler(t *testing.T) {
path := filepath.Join(t.TempDir(), "config.toml")
if err := os.WriteFile(path, []byte("workers = 2\n\n[[users]]\nid = 42\n"), 0o644); err != nil {
t.Fatal(err)
}
if err := config.Init(t.Context(), path); err != nil {
t.Fatal(err)
}
update := &ext.Update{CallbackQuery: &tg.UpdateBotCallbackQuery{UserID: 42}}
called := false
handler := withPermission(func(ctx *ext.Context, u *ext.Update) error {
called = true
return nil
})
if err := handler(&ext.Context{}, update); err != nil {
t.Fatalf("withPermission returned error: %v", err)
}
if !called {
t.Fatal("withPermission did not invoke the wrapped handler")
}
}

View File

@@ -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"), 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.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.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)))

View File

@@ -13,7 +13,6 @@ import (
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/i18n/i18nk"
"github.com/krau/SaveAny-Bot/common/utils/strutil"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/database"
"github.com/krau/SaveAny-Bot/pkg/rule"
)
@@ -85,46 +84,6 @@ func handleRuleCmd(ctx *ext.Context, update *ext.Update) error {
return dispatcher.EndGroups
}
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgRuleInfoCreateRuleSuccess, nil)), nil)
case "preset":
// /rule preset <storage> [base_path]
if len(args) < 3 {
ctx.Reply(update, ext.ReplyTextStyledTextArray(msgelem.BuildRuleHelpStyling(user.ApplyRule, user.Rules)), nil)
return dispatcher.EndGroups
}
storageName := args[2]
if !config.C().HasStorage(user.ChatID, storageName) {
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgRuleErrorStorageNotFound, map[string]any{
"Storage": storageName,
})), nil)
return dispatcher.EndGroups
}
basePath := ""
if len(args) >= 4 {
basePath = args[3]
}
presets := rule.PresetCategories(basePath)
imported := 0
for _, p := range presets {
rd := &database.Rule{
Type: rule.FileNameRegex.String(),
Data: p.Regex,
StorageName: storageName,
DirPath: p.Dir,
UserID: user.ID,
}
if err := database.CreateRule(ctx, rd); err != nil {
logger.Errorf("failed to create preset rule %s: %s", p.Name, err)
continue
}
imported++
}
if imported == 0 {
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgRuleErrorCreateRuleFailed, nil)), nil)
return dispatcher.EndGroups
}
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgRuleInfoPresetImported, map[string]any{
"Count": imported,
})), nil)
case "del":
// /rule del <id>
if len(args) < 3 {

View File

@@ -1,7 +1,6 @@
package handlers
import (
"fmt"
"strings"
"github.com/celestix/gotgproto/dispatcher"
@@ -44,11 +43,7 @@ func handleSilentCmd(ctx *ext.Context, update *ext.Update) error {
}
func handleSetDefaultCallback(ctx *ext.Context, update *ext.Update) error {
dataParts := strings.Split(string(update.CallbackQuery.Data), " ")
if len(dataParts) < 2 {
return fmt.Errorf("invalid callback data: %q", update.CallbackQuery.Data)
}
dataid := dataParts[1]
dataid := strings.Split(string(update.CallbackQuery.Data), " ")[1]
data, ok := cache.Get[tcbdata.SetDefaultStorage](dataid)
failedAnswer := func(message string) error {

View File

@@ -89,11 +89,7 @@ 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)})),
)
const maxShown = 10
for i, t := range tasks {
if i >= maxShown {
break
}
for _, t := range tasks {
created := t.Created.In(time.Local).Format("2006-01-02 15:04:05")
status := i18n.T(i18nk.BotMsgTasksStatusQueued)
if t.Cancelled {
@@ -109,9 +105,10 @@ func showQueuedTasks(ctx *ext.Context, update *ext.Update) {
styling.Plain("\n"+i18n.T(i18nk.BotMsgTasksFieldStatus)),
styling.Code(status),
)
}
if len(tasks) > maxShown {
opts = append(opts, styling.Plain("\n"+i18n.T(i18nk.BotMsgTasksTruncatedNote, map[string]any{"Count": len(tasks)})))
if len(tasks) > 10 {
opts = append(opts, styling.Plain("\n"+i18n.T(i18nk.BotMsgTasksTruncatedNote, map[string]any{"Count": len(tasks)})))
break
}
}
ctx.Reply(update, ext.ReplyTextStyledTextArray(opts), nil)
}

View File

@@ -10,7 +10,6 @@ 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"
@@ -101,10 +100,7 @@ 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 {
ctx.AnswerCallback(msgelem.AlertCallbackAnswer(u.CallbackQuery.GetQueryID(), i18n.T(i18nk.BotMsgUpdateErrorVersionVarInvalid, map[string]any{
"Error": err.Error(),
})))
return dispatcher.EndGroups
return err
}
ctx.EditMessage(u.GetUserChat().GetID(), &tg.MessagesEditMessageRequest{
ID: u.CallbackQuery.GetMsgID(),

View File

@@ -24,8 +24,6 @@ func BuildRuleHelpStyling(enabled bool, rules []database.Rule) []styling.StyledT
styling.Plain(i18n.T(i18nk.BotMsgRuleHelpSwitchSuffix, nil)),
styling.Code("add"),
styling.Plain(i18n.T(i18nk.BotMsgRuleHelpAddSuffix, nil)),
styling.Code("preset"),
styling.Plain(i18n.T(i18nk.BotMsgRuleHelpPresetSuffix, nil)),
styling.Code("del"),
styling.Plain(i18n.T(i18nk.BotMsgRuleHelpDelSuffix, nil)),
styling.Plain(i18n.T(i18nk.BotMsgRuleHelpExistingRulesPrefix, nil)),

View File

@@ -3,8 +3,8 @@ package re
import "regexp"
var (
TgMessageLinkRegexString = `https?://(?:t|telegram)\.me/(?:c/\d+|[A-Za-z0-9_]+)/\d+(?:/\d+)?(?:\?[^\s#]*[A-Za-z0-9_])?\b`
TgMessageLinkRegexString = `https?://t\.me/(?:c/\d+|[A-Za-z0-9_]+)/\d+(?:/\d+)?(?:\?[^\s#]*[A-Za-z0-9_])?\b`
TgMessageLinkRegexp = regexp.MustCompile(TgMessageLinkRegexString)
TelegraphUrlRegexString = `https://telegra\.ph/[^\s]+`
TelegraphUrlRegexString = `https://telegra.ph/.*`
TelegraphUrlRegexp = regexp.MustCompile(TelegraphUrlRegexString)
)

View File

@@ -3,7 +3,6 @@ package shortcut
import (
"encoding/json"
"fmt"
"net/url"
"strings"
@@ -180,19 +179,21 @@ type TelegraphResult struct {
// return replied message, image urls, telegraph path(unescaped), error
func GetTphPicsFromMessageWithReply(ctx *ext.Context, update *ext.Update) (*types.Message, *TelegraphResult, error) {
logger := log.FromContext(ctx)
tphurl := findTelegraphURL(update.EffectiveMessage.Message)
tphurl := re.TelegraphUrlRegexp.FindString(tgutil.ExtractMessageEntityUrlsText(update.EffectiveMessage.Message))
if tphurl == "" {
logger.Warnf("No telegraph url found but called handleTelegraph")
return nil, nil, dispatcher.ContinueGroups
}
pagepath, err := parseTelegraphPagePath(tphurl)
pagepath := strings.Split(tphurl, "/")[len(strings.Split(tphurl, "/"))-1]
tphdir, err := url.PathUnescape(pagepath)
if err != nil {
logger.Errorf("Failed to parse telegraph path: %s", err)
logger.Errorf("Failed to unescape telegraph path: %s", err)
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgCommonErrorParseTelegraphPathFailed, map[string]any{
"Error": err.Error(),
})), nil)
return nil, nil, dispatcher.EndGroups
}
tphdir = strings.TrimSpace(tphdir)
msg, err := ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgCommonInfoFetchingTelegraphPage, nil)), nil)
if err != nil {
logger.Errorf("Failed to reply to update: %s", err)
@@ -243,57 +244,7 @@ func GetTphPicsFromMessageWithReply(ctx *ext.Context, update *ext.Update) (*type
}
return msg, &TelegraphResult{
Pics: imgs,
TphDir: pagepath,
TphDir: tphdir,
Page: page,
}, nil
}
func findTelegraphURL(msg *tg.Message) string {
if msg == nil {
return ""
}
var firstMatch string
findValid := func(text string) string {
for _, tphurl := range re.TelegraphUrlRegexp.FindAllString(text, -1) {
if firstMatch == "" {
firstMatch = tphurl
}
if _, err := parseTelegraphPagePath(tphurl); err == nil {
return tphurl
}
}
return ""
}
for _, entityURL := range tgutil.ExtractMessageEntityUrls(msg) {
if tphurl := findValid(entityURL); tphurl != "" {
return tphurl
}
}
if tphurl := findValid(msg.GetMessage()); tphurl != "" {
return tphurl
}
return firstMatch
}
func parseTelegraphPagePath(pageURL string) (string, error) {
u, err := url.Parse(pageURL)
if err != nil {
return "", fmt.Errorf("invalid telegraph URL: %w", err)
}
if u.Scheme != "https" || !strings.EqualFold(u.Hostname(), "telegra.ph") {
return "", fmt.Errorf("invalid telegraph URL host: %s", u.Host)
}
pagepath := strings.Trim(u.EscapedPath(), "/")
if pagepath == "" || strings.Contains(pagepath, "/") {
return "", fmt.Errorf("invalid telegraph URL path: %s", u.Path)
}
pagepath, err = url.PathUnescape(pagepath)
if err != nil {
return "", fmt.Errorf("failed to unescape telegraph path: %w", err)
}
pagepath = strings.TrimSpace(pagepath)
if pagepath == "" || strings.Contains(pagepath, "/") {
return "", fmt.Errorf("invalid telegraph URL path: %s", u.Path)
}
return pagepath, nil
}

View File

@@ -1,163 +0,0 @@
package shortcut
import (
"testing"
"github.com/gotd/td/tg"
)
func TestFindTelegraphURL(t *testing.T) {
tests := []struct {
name string
msg *tg.Message
want string
}{
{
name: "single URL entity",
msg: &tg.Message{
Message: "https://telegra.ph/Example-01-02",
Entities: []tg.MessageEntityClass{
&tg.MessageEntityURL{Offset: 0, Length: 32},
},
},
want: "https://telegra.ph/Example-01-02",
},
{
name: "Telegraph URL before another URL",
msg: &tg.Message{
Message: "https://telegra.ph/Example-01-02 https://example.com/",
Entities: []tg.MessageEntityClass{
&tg.MessageEntityURL{Offset: 0, Length: 32},
&tg.MessageEntityURL{Offset: 33, Length: 20},
},
},
want: "https://telegra.ph/Example-01-02",
},
{
name: "hidden Telegraph URL",
msg: &tg.Message{
Message: "article",
Entities: []tg.MessageEntityClass{
&tg.MessageEntityTextURL{
Offset: 0,
Length: 7,
URL: "https://telegra.ph/Hidden-01-02",
},
},
},
want: "https://telegra.ph/Hidden-01-02",
},
{
name: "URL entity after non-BMP character",
msg: &tg.Message{
Message: "😀 https://telegra.ph/Emoji-01-02",
Entities: []tg.MessageEntityClass{
&tg.MessageEntityURL{Offset: 3, Length: 30},
},
},
want: "https://telegra.ph/Emoji-01-02",
},
{
name: "valid Telegraph URL after invalid candidate",
msg: &tg.Message{
Message: "https://telegra.ph/nested/Bad https://telegra.ph/Valid-01-02",
Entities: []tg.MessageEntityClass{
&tg.MessageEntityURL{Offset: 0, Length: 29},
&tg.MessageEntityURL{Offset: 30, Length: 30},
},
},
want: "https://telegra.ph/Valid-01-02",
},
{
name: "plain message fallback",
msg: &tg.Message{Message: "read https://telegra.ph/Plain-01-02 now"},
want: "https://telegra.ph/Plain-01-02",
},
{
name: "nil message",
msg: nil,
want: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := findTelegraphURL(tt.msg); got != tt.want {
t.Fatalf("findTelegraphURL() = %q, want %q", got, tt.want)
}
})
}
}
func TestParseTelegraphPagePath(t *testing.T) {
tests := []struct {
name string
pageURL string
want string
wantErr bool
}{
{
name: "plain path",
pageURL: "https://telegra.ph/Example-01-02",
want: "Example-01-02",
},
{
name: "escaped path with query and fragment",
pageURL: "https://telegra.ph/%E6%B5%8B%E8%AF%95-01-02?source=telegram#top",
want: "测试-01-02",
},
{
name: "trailing slash",
pageURL: "https://telegra.ph/Example-01-02/",
want: "Example-01-02",
},
{
name: "root URL",
pageURL: "https://telegra.ph/",
wantErr: true,
},
{
name: "wrong host",
pageURL: "https://example.com/Example-01-02",
wantErr: true,
},
{
name: "wrong scheme",
pageURL: "http://telegra.ph/Example-01-02",
wantErr: true,
},
{
name: "invalid percent escape",
pageURL: "https://telegra.ph/Invalid-%zz",
wantErr: true,
},
{
name: "nested path",
pageURL: "https://telegra.ph/nested/Example-01-02",
wantErr: true,
},
{
name: "encoded slash",
pageURL: "https://telegra.ph/nested%2FExample-01-02",
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := parseTelegraphPagePath(tt.pageURL)
if tt.wantErr {
if err == nil {
t.Fatalf("parseTelegraphPagePath(%q) returned no error", tt.pageURL)
}
return
}
if err != nil {
t.Fatalf("parseTelegraphPagePath(%q) failed: %v", tt.pageURL, err)
}
if got != tt.want {
t.Fatalf("parseTelegraphPagePath(%q) = %q, want %q", tt.pageURL, got, tt.want)
}
})
}
}

View File

@@ -64,7 +64,7 @@ func handleWatchCmd(ctx *ext.Context, update *ext.Update) error {
filter := ""
if len(args) > 2 {
filterArg := strings.Join(args[2:], " ")
filterType, _, _ := strings.Cut(filterArg, ":")
filterType := strings.Split(filterArg, ":")[0]
filterData := strings.Split(filterArg, ":")[1]
if filterType == "" || filterData == "" {
ctx.Reply(update, ext.ReplyTextString(i18n.T(i18nk.BotMsgWatchErrorFilterFormatInvalid)), nil)

View File

@@ -2,7 +2,6 @@ package user
import (
"context"
"sync"
"time"
"github.com/celestix/gotgproto"
@@ -21,18 +20,17 @@ import (
)
var uc *gotgproto.Client
// 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()
})
var ectx *ext.Context
func GetCtx() *ext.Context {
if ectx != nil {
return ectx
}
if uc == nil {
return nil
}
return getEctx()
ectx = uc.CreateContext()
return ectx
}
func Login(ctx context.Context) (*gotgproto.Client, error) {

View File

@@ -22,13 +22,7 @@ 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 {
@@ -48,12 +42,7 @@ func main() {
return fmt.Errorf("failed to parse yaml %s: %w", path, err)
}
fileKeys := make(map[string]struct{})
collectKeys(content, "", fileKeys)
localeFiles = append(localeFiles, localeFile{path: path, keys: fileKeys})
for k := range fileKeys {
keys[k] = struct{}{}
}
collectKeys(content, "", keys)
return nil
})
if err != nil {
@@ -61,25 +50,6 @@ 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)

View File

@@ -10,14 +10,12 @@ import (
"slices"
"github.com/charmbracelet/log"
"github.com/gotd/td/telegram/downloader"
"github.com/krau/SaveAny-Bot/api"
"github.com/krau/SaveAny-Bot/client/bot"
userclient "github.com/krau/SaveAny-Bot/client/user"
"github.com/krau/SaveAny-Bot/common/cache"
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/common/utils/fsutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/core"
"github.com/krau/SaveAny-Bot/database"
@@ -58,25 +56,11 @@ func Run(cmd *cobra.Command, _ []string) {
cancel()
}()
core.SetDownloaderProvider(func() downloader.Client {
if ectx := bot.ExtContext(); ectx != nil {
return ectx.Raw
}
return nil
})
// 恢复任务携带 ext 上下文, 让进度编辑/取消按钮在恢复后继续工作。
recoverCtx := context.Background()
if ectx := bot.ExtContext(); ectx != nil {
recoverCtx = tgutil.ExtWithContext(recoverCtx, ectx)
}
core.RecoverTasks(recoverCtx)
core.Run(ctx)
<-ctx.Done()
logger.Info("Exiting...")
defer logger.Info("Exit complete")
core.Close()
cleanCache()
}
@@ -103,7 +87,7 @@ func initAll(ctx context.Context) (<-chan struct{}, error) {
}
}
if err := api.Start(ctx); err != nil {
logger.Fatal("Failed to start API server", "error", err)
logger.Error("Failed to start API server", "error", err)
}
return bot.Init(ctx), nil
}
@@ -117,15 +101,6 @@ func cleanCache() {
log.Error("Invalid cache directory", "path", config.C().Temp.BasePath)
return
}
unfinished, err := database.CountUnfinishedTasks(context.Background())
if err != nil {
log.Error("Failed to count unfinished tasks, skipping cache cleanup", "error", err)
return
}
if unfinished > 0 {
log.Info("Skipping cache cleanup: unfinished tasks need their cache files for recovery", "tasks", unfinished)
return
}
currentDir, err := os.Getwd()
if err != nil {
log.Error("Failed to get working directory", "error", err)

View File

@@ -2,7 +2,6 @@ package cache
import (
"fmt"
"sync"
"time"
"github.com/charmbracelet/log"
@@ -10,26 +9,24 @@ import (
"github.com/krau/SaveAny-Bot/config"
)
var (
cache *ristretto.Cache[string, any]
initOnce sync.Once
)
var cache *ristretto.Cache[string, any]
func Init() {
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 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)
},
})
if err != nil {
log.Fatalf("failed to create ristretto cache: %v", err)
}
cache = c
}
func Set(key string, value any) error {

View File

@@ -84,8 +84,8 @@ const (
BotMsgCommonPromptSelectDefaultDir Key = "bot.msg.common.prompt_select_default_dir"
BotMsgCommonPromptSelectDefaultStorage Key = "bot.msg.common.prompt_select_default_storage"
BotMsgCommonPromptSelectDir Key = "bot.msg.common.prompt_select_dir"
BotMsgConfigButtonConflictStrategy Key = "bot.msg.config.button_conflict_strategy"
BotMsgConfigButtonFilenameStrategy Key = "bot.msg.config.button_filename_strategy"
BotMsgConfigButtonConflictStrategy Key = "bot.msg.config.button_conflict_strategy"
BotMsgConfigConflictStrategyAsk Key = "bot.msg.config.conflict_strategy_ask"
BotMsgConfigConflictStrategyOverwrite Key = "bot.msg.config.conflict_strategy_overwrite"
BotMsgConfigConflictStrategyRename Key = "bot.msg.config.conflict_strategy_rename"
@@ -93,8 +93,8 @@ const (
BotMsgConfigErrorInvalidCallbackData Key = "bot.msg.config.error_invalid_callback_data"
BotMsgConfigErrorInvalidTemplate Key = "bot.msg.config.error_invalid_template"
BotMsgConfigFnametmplHelp Key = "bot.msg.config.fnametmpl_help"
BotMsgConfigInfoConflictStrategySet Key = "bot.msg.config.info_conflict_strategy_set"
BotMsgConfigInfoCurrentTemplatePrefix Key = "bot.msg.config.info_current_template_prefix"
BotMsgConfigInfoConflictStrategySet Key = "bot.msg.config.info_conflict_strategy_set"
BotMsgConfigInfoFilenameStrategySet Key = "bot.msg.config.info_filename_strategy_set"
BotMsgConfigInfoTemplateUpdated Key = "bot.msg.config.info_template_updated"
BotMsgConfigPromptSelectConflictStrategy Key = "bot.msg.config.prompt_select_conflict_strategy"
@@ -150,52 +150,28 @@ const (
BotMsgProgressAria2Downloading Key = "bot.msg.progress.aria2_downloading"
BotMsgProgressAria2Start Key = "bot.msg.progress.aria2_start"
BotMsgProgressAvgSpeedPrefix Key = "bot.msg.progress.avg_speed_prefix"
BotMsgProgressBatchCanceled Key = "bot.msg.progress.batch_canceled"
BotMsgProgressBatchDone Key = "bot.msg.progress.batch_done"
BotMsgProgressBatchDoneWithSkipped Key = "bot.msg.progress.batch_done_with_skipped"
BotMsgProgressBatchFailedGroup Key = "bot.msg.progress.batch_failed_group"
BotMsgProgressBatchFailedItem Key = "bot.msg.progress.batch_failed_item"
BotMsgProgressBatchFailedTask Key = "bot.msg.progress.batch_failed_task"
BotMsgProgressBatchFailureStageBatchUpload Key = "bot.msg.progress.batch_failure_stage_batch_upload"
BotMsgProgressBatchFailureStageCache Key = "bot.msg.progress.batch_failure_stage_cache"
BotMsgProgressBatchFailureStageConfirm Key = "bot.msg.progress.batch_failure_stage_confirm"
BotMsgProgressBatchFailureStageDownload Key = "bot.msg.progress.batch_failure_stage_download"
BotMsgProgressBatchFailureStageInternal Key = "bot.msg.progress.batch_failure_stage_internal"
BotMsgProgressBatchFailureStageUpload Key = "bot.msg.progress.batch_failure_stage_upload"
BotMsgProgressBatchItemConfirming Key = "bot.msg.progress.batch_item_confirming"
BotMsgProgressBatchItemDownloading Key = "bot.msg.progress.batch_item_downloading"
BotMsgProgressBatchItemDownloadingUnknown Key = "bot.msg.progress.batch_item_downloading_unknown"
BotMsgProgressBatchItemRetrying Key = "bot.msg.progress.batch_item_retrying"
BotMsgProgressBatchItemTransferring Key = "bot.msg.progress.batch_item_transferring"
BotMsgProgressBatchItemTransferringUnknown Key = "bot.msg.progress.batch_item_transferring_unknown"
BotMsgProgressBatchItemUploading Key = "bot.msg.progress.batch_item_uploading"
BotMsgProgressBatchStatusHeader Key = "bot.msg.progress.batch_status_header"
BotMsgProgressBatchSummaryConfirming Key = "bot.msg.progress.batch_summary_confirming"
BotMsgProgressBatchSummaryFailed Key = "bot.msg.progress.batch_summary_failed"
BotMsgProgressBatchSummaryHiddenActive Key = "bot.msg.progress.batch_summary_hidden_active"
BotMsgProgressBatchSummarySkipped Key = "bot.msg.progress.batch_summary_skipped"
BotMsgProgressBatchDonePrefix Key = "bot.msg.progress.batch_done_prefix"
BotMsgProgressBatchProcessingPrefix Key = "bot.msg.progress.batch_processing_prefix"
BotMsgProgressBatchStartPrefix Key = "bot.msg.progress.batch_start_prefix"
BotMsgProgressCurrentProgressPrefix Key = "bot.msg.progress.current_progress_prefix"
BotMsgProgressCurrentSpeedPrefix Key = "bot.msg.progress.current_speed_prefix"
BotMsgProgressDirectDonePrefix Key = "bot.msg.progress.direct_done_prefix"
BotMsgProgressDirectStart Key = "bot.msg.progress.direct_start"
BotMsgProgressDownloadDonePrefix Key = "bot.msg.progress.download_done_prefix"
BotMsgProgressDownloadFailedPrefix Key = "bot.msg.progress.download_failed_prefix"
BotMsgProgressDownloadedPrefix Key = "bot.msg.progress.downloaded_prefix"
BotMsgProgressDownloadingPrefix Key = "bot.msg.progress.downloading_prefix"
BotMsgProgressErrorPrefix Key = "bot.msg.progress.error_prefix"
BotMsgProgressFileNamePrefix Key = "bot.msg.progress.file_name_prefix"
BotMsgProgressFileProcessingPrefix Key = "bot.msg.progress.file_processing_prefix"
BotMsgProgressFileSizePrefix Key = "bot.msg.progress.file_size_prefix"
BotMsgProgressFileStartPrefix Key = "bot.msg.progress.file_start_prefix"
BotMsgProgressParsedDonePrefix Key = "bot.msg.progress.parsed_done_prefix"
BotMsgProgressParsedStartPrefix Key = "bot.msg.progress.parsed_start_prefix"
BotMsgProgressProcessingListPrefix Key = "bot.msg.progress.processing_list_prefix"
BotMsgProgressProcessingNone Key = "bot.msg.progress.processing_none"
BotMsgProgressSavePathPrefix Key = "bot.msg.progress.save_path_prefix"
BotMsgProgressSingleCanceled Key = "bot.msg.progress.single_canceled"
BotMsgProgressSingleDone Key = "bot.msg.progress.single_done"
BotMsgProgressSingleDownloading Key = "bot.msg.progress.single_downloading"
BotMsgProgressSingleDownloadingUnknown Key = "bot.msg.progress.single_downloading_unknown"
BotMsgProgressSingleFailed Key = "bot.msg.progress.single_failed"
BotMsgProgressSingleStatusHeader Key = "bot.msg.progress.single_status_header"
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"
BotMsgProgressTaskCanceled Key = "bot.msg.progress.task_canceled"
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"
@@ -224,7 +200,6 @@ const (
BotMsgRuleErrorGetUserRulesFailed Key = "bot.msg.rule.error_get_user_rules_failed"
BotMsgRuleErrorInvalidRuleId Key = "bot.msg.rule.error_invalid_rule_id"
BotMsgRuleErrorInvalidRuleType Key = "bot.msg.rule.error_invalid_rule_type"
BotMsgRuleErrorStorageNotFound Key = "bot.msg.rule.error_storage_not_found"
BotMsgRuleErrorUpdateUserFailed Key = "bot.msg.rule.error_update_user_failed"
BotMsgRuleHelpAddSuffix Key = "bot.msg.rule.help_add_suffix"
BotMsgRuleHelpAvailableOps Key = "bot.msg.rule.help_available_ops"
@@ -232,20 +207,18 @@ const (
BotMsgRuleHelpCurrentModeEnabled Key = "bot.msg.rule.help_current_mode_enabled"
BotMsgRuleHelpDelSuffix Key = "bot.msg.rule.help_del_suffix"
BotMsgRuleHelpExistingRulesPrefix Key = "bot.msg.rule.help_existing_rules_prefix"
BotMsgRuleHelpPresetSuffix Key = "bot.msg.rule.help_preset_suffix"
BotMsgRuleHelpSwitchSuffix Key = "bot.msg.rule.help_switch_suffix"
BotMsgRuleHelpUsage Key = "bot.msg.rule.help_usage"
BotMsgRuleInfoCreateRuleSuccess Key = "bot.msg.rule.info_create_rule_success"
BotMsgRuleInfoDeleteRuleSuccess Key = "bot.msg.rule.info_delete_rule_success"
BotMsgRuleInfoPresetImported Key = "bot.msg.rule.info_preset_imported"
BotMsgRuleInfoRuleModeDisabled Key = "bot.msg.rule.info_rule_mode_disabled"
BotMsgRuleInfoRuleModeEnabled Key = "bot.msg.rule.info_rule_mode_enabled"
BotMsgRulePromptProvideRuleId Key = "bot.msg.rule.prompt_provide_rule_id"
BotMsgRulePromptProvideStorageName Key = "bot.msg.rule.prompt_provide_storage_name"
BotMsgSaveErrorInvalidIdOrUsername Key = "bot.msg.save.error_invalid_id_or_username"
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"

View File

@@ -196,11 +196,7 @@ bot:
help_switch_suffix: " - Toggle rule mode\n"
help_add_suffix: " <type> <data> <storage_name> <path> - Add rule\n"
help_del_suffix: " <rule_id> - Delete rule\n"
help_preset_suffix: " <storage_name> [base_path] - Import built-in filetype rules (video/image/audio/document/archive)\n"
help_existing_rules_prefix: "\nCurrent rules:\n"
prompt_provide_storage_name: "Please provide a storage name"
error_storage_not_found: "Storage not found: {{.Storage}}"
info_preset_imported: "Imported {{.Count}} built-in classification rules into storage {{.Storage}}"
dir:
error_get_user_dirs_failed: "Failed to get user directories"
error_get_user_failed: "Failed to get user"
@@ -351,56 +347,32 @@ bot:
info_filename_prefix: "Filename: "
info_prompt_select_storage: "\nPlease select storage"
progress:
batch_status_header: "<b>📦 Processing</b>\n\nFiles: <code>{{.Total}}</code> | Total size: <code>{{.TotalSize}}</code>\nStatus: ✅ <code>{{.Completed}}</code> | 📥 <code>{{.Downloaded}}</code> | ⏳ <code>{{.Waiting}}</code>\nTotal speed: ⬇️ <code>{{.DownloadSpeed}}</code> | ⬆️ <code>{{.UploadSpeed}}</code>"
batch_item_downloading: "<blockquote><b>⬇️ {{.Index}}/{{.Total}} Downloading</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / <code>{{.Size}}</code></blockquote>"
batch_item_downloading_unknown: "<blockquote><b>⬇️ {{.Index}}/{{.Total}} Downloading</b>\n<code>{{.Name}}</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / unknown</blockquote>"
batch_item_transferring: "<blockquote><b>↕️ {{.Index}}/{{.Total}} Transferring</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / <code>{{.Size}}</code></blockquote>"
batch_item_transferring_unknown: "<blockquote><b>↕️ {{.Index}}/{{.Total}} Transferring</b>\n<code>{{.Name}}</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / unknown</blockquote>"
batch_item_uploading: "<blockquote><b>⬆️ {{.Index}}/{{.Total}} Uploading</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / <code>{{.Size}}</code></blockquote>"
batch_item_retrying: "<blockquote><b>🔁 {{.Index}}/{{.Total}} Retrying upload</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\nRetry: <code>{{.Attempt}}/{{.Limit}}</code>\nSpeed before failure: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / <code>{{.Size}}</code>\nReason: <code>{{.Reason}}</code></blockquote>"
batch_item_confirming: "<blockquote><b>⏳ {{.Index}}/{{.Total}} Waiting</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n<i>Uploaded, awaiting remote confirmation</i></blockquote>"
batch_summary_hidden_active: "🔄 <code>{{.Count}}</code> more files are active"
batch_summary_confirming: "☁️ Uploaded, awaiting group send: <code>{{.Count}}</code>"
batch_summary_failed: "❌ Failed: <code>{{.Count}}</code>"
batch_summary_skipped: "⏭️ Skipped: <code>{{.Count}}</code>"
batch_done: "<b>✅ Completed</b>\n\nFiles: <code>{{.Count}}</code>\nTotal size: <code>{{.Size}}</code>"
batch_done_with_skipped: "<b>⚠️ Completed</b>\n\nSucceeded: <code>{{.Success}}</code>\nSkipped: <code>{{.Skipped}}</code>\nTotal size: <code>{{.Size}}</code>"
batch_canceled: "<b>🚫 Task canceled</b>\n\nFiles: <code>{{.Total}}</code>\nCompleted: <code>{{.Completed}}</code>\nIncomplete: <code>{{.Incomplete}}</code>\nSkipped: <code>{{.Skipped}}</code>"
batch_failed_item: "<b>❌ Processing failed</b>\n\nFailed file: <code>{{.Index}}. {{.Name}}</code>\nStage: <code>{{.Stage}}</code>\nProgress: <code>{{.Progress}}</code>\nSpeed before failure: <code>{{.Speed}}</code>\nReason: <code>{{.Reason}}</code>\n\n✅ Completed: <code>{{.Completed}}</code>\n❌ Failed: <code>{{.Failed}}</code>\n⏹ Incomplete: <code>{{.Incomplete}}</code>"
batch_failed_group: "<b>❌ Batch upload failed</b>\n\nAffected files: <code>{{.Affected}}</code>\nReason: <code>{{.Reason}}</code>\n\n✅ Completed: <code>{{.Completed}}</code>\n❌ Batch failed: <code>{{.Failed}}</code>\n⏹ Incomplete: <code>{{.Incomplete}}</code>"
batch_failed_task: "<b>❌ Processing failed</b>\n\nReason: <code>{{.Reason}}</code>\n\n✅ Completed: <code>{{.Completed}}</code>\n⏹ Incomplete: <code>{{.Incomplete}}</code>"
batch_failure_stage_download: "download"
batch_failure_stage_cache: "local cache"
batch_failure_stage_upload: "upload"
batch_failure_stage_confirm: "remote confirmation"
batch_failure_stage_batch_upload: "batch upload"
batch_failure_stage_internal: "internal task"
single_status_header: "<b>📦 Processing</b>"
single_downloading: "<blockquote><b>⬇️ Downloading</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / <code>{{.Size}}</code>\nSave to: <code>{{.Destination}}</code></blockquote>"
single_downloading_unknown: "<blockquote><b>⬇️ Downloading</b>\n<code>{{.Name}}</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / unknown\nSave to: <code>{{.Destination}}</code></blockquote>"
single_uploading: "<blockquote><b>⬆️ Uploading</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / <code>{{.Size}}</code>\nSave to: <code>{{.Destination}}</code></blockquote>"
single_upload_retrying: "<blockquote><b>🔁 Retrying upload</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\nAttempt: <code>{{.Attempt}}</code>\nSpeed: <code>{{.Speed}}</code>\nSize: <code>{{.Current}}</code> / <code>{{.Size}}</code>\nSave to: <code>{{.Destination}}</code></blockquote>"
single_done: "<b>✅ Completed</b>\n\nFilename: <code>{{.Name}}</code>\nTotal size: <code>{{.Size}}</code>\nSave to: <code>{{.Destination}}</code>"
single_canceled: "<b>🚫 Task canceled</b>\n\nFilename: <code>{{.Name}}</code>"
single_failed: "<b>❌ Processing failed</b>\n\nFilename: <code>{{.Name}}</code>\nReason: <code>{{.Reason}}</code>"
batch_start_prefix: "Starting batch download task\nTotal size: "
batch_processing_prefix: "Processing batch download task\nTotal size: "
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: "
current_progress_prefix: "\nCurrent progress: "
task_canceled: "Task canceled"
task_canceled_with_id: "Processing canceled: {{.TaskID}}"
task_failed_with_error: "Processing failed: {{.Error}}"
batch_done_prefix: "Completed\nFile count: "
direct_done_prefix: "Completed, file count: "
parsed_start_prefix: "Starting download from {{.Site}}\nTotal size: "
parsed_done_prefix: "Completed, resource count: "
telegraph_start_prefix: "Starting Telegraph download\nImage count: "
telegraph_progress_prefix: "Downloading\nCurrent progress: "
telegraph_done_prefix: "Completed\nImage count: "
file_start_prefix: "Starting download\nFilename: "
file_processing_prefix: "Processing download task\nFilename: "
download_failed_prefix: "Download failed\nFilename: "
download_done_prefix: "Download completed\nFilename: "
file_size_prefix: "\nFile size: "
save_path_prefix: "\nSave path: "
total_size_prefix: "\nTotal size: "
direct_start: "Starting download, total size: {{.SizeMB}} MB ({{.Count}} files)"
file_name_prefix: "Filename: "
error_prefix: "\nError: "
aria2_start: "Waiting for Aria2 to complete download (GID: {{.GID}})..."
aria2_downloading: "Aria2 downloading (GID: {{.GID}})\n"
@@ -426,7 +398,7 @@ bot:
transfer_failed_files_prefix: "\nFailed files: "
syncpeers:
start: "Starting to sync peers..."
success: "Peer sync completed, total {{.Count}} chats synced"
done: "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"

View File

@@ -197,11 +197,7 @@ bot:
help_switch_suffix: " - 开关规则模式\n"
help_add_suffix: " <类型> <数据> <存储名> <路径> - 添加规则\n"
help_del_suffix: " <规则ID> - 删除规则\n"
help_preset_suffix: " <存储名> [基础路径] - 导入内置文件类型分类规则(视频/图片/音频/文档/压缩包)\n"
help_existing_rules_prefix: "\n当前已添加的规则:\n"
prompt_provide_storage_name: "请提供存储名称"
error_storage_not_found: "未找到存储: {{.Storage}}"
info_preset_imported: "已导入 {{.Count}} 条内置分类规则到存储 {{.Storage}}"
dir:
error_get_user_dirs_failed: "获取用户文件夹失败"
error_get_user_failed: "获取用户失败"
@@ -238,9 +234,9 @@ bot:
info_install_plugin_success: "插件安装成功: {{.Name}}"
parse:
info_parsing: "正在解析..."
error_parse_text_failed: "解析文本失败: {{.Error}}"
error_build_storage_select_keyboard_failed: "构建存储选择键盘失败: {{.Error}}"
error_build_parsed_text_entity_failed: "构建解析文本实体失败: {{.Error}}"
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}}"
info_link_prefix: "\n链接: "
info_author_prefix: "\n作者: "
info_description_prefix: "\n描述: "
@@ -352,56 +348,32 @@ bot:
info_filename_prefix: "文件名: "
info_prompt_select_storage: "\n请选择存储位置"
progress:
batch_status_header: "<b>📦 正在处理</b>\n\n文件<code>{{.Total}}</code> 总大小:<code>{{.TotalSize}}</code>\n状态✅ <code>{{.Completed}}</code> 📥 <code>{{.Downloaded}}</code> ⏳ <code>{{.Waiting}}</code>\n总速度 <code>{{.DownloadSpeed}}</code> ⬆️ <code>{{.UploadSpeed}}</code>"
batch_item_downloading: "<blockquote><b>⬇️ {{.Index}}/{{.Total}} 下载中</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / <code>{{.Size}}</code></blockquote>"
batch_item_downloading_unknown: "<blockquote><b>⬇️ {{.Index}}/{{.Total}} 下载中</b>\n<code>{{.Name}}</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / 未知</blockquote>"
batch_item_transferring: "<blockquote><b>↕️ {{.Index}}/{{.Total}} 传输中</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / <code>{{.Size}}</code></blockquote>"
batch_item_transferring_unknown: "<blockquote><b>↕️ {{.Index}}/{{.Total}} 传输中</b>\n<code>{{.Name}}</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / 未知</blockquote>"
batch_item_uploading: "<blockquote><b>⬆️ {{.Index}}/{{.Total}} 上传中</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / <code>{{.Size}}</code></blockquote>"
batch_item_retrying: "<blockquote><b>🔁 {{.Index}}/{{.Total}} 上传重试</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n重试次数<code>{{.Attempt}}/{{.Limit}}</code>\n失败前速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / <code>{{.Size}}</code>\n原因<code>{{.Reason}}</code></blockquote>"
batch_item_confirming: "<blockquote><b>⏳ {{.Index}}/{{.Total}} 等待中</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n<i>文件已上传,正在等待远端确认</i></blockquote>"
batch_summary_hidden_active: "🔄 另有 <code>{{.Count}}</code> 个文件正在处理"
batch_summary_confirming: "☁️ 已上传,等待整组发送:<code>{{.Count}}</code>"
batch_summary_failed: "❌ 失败:<code>{{.Count}}</code>"
batch_summary_skipped: "⏭️ 已跳过:<code>{{.Count}}</code>"
batch_done: "<b>✅ 处理完成</b>\n\n文件数: <code>{{.Count}}</code>\n总大小: <code>{{.Size}}</code>"
batch_done_with_skipped: "<b>⚠️ 处理完成</b>\n\n成功: <code>{{.Success}}</code>\n已跳过: <code>{{.Skipped}}</code>\n总大小: <code>{{.Size}}</code>"
batch_canceled: "<b>🚫 任务已取消</b>\n\n文件数: <code>{{.Total}}</code>\n已完成: <code>{{.Completed}}</code>\n未完成: <code>{{.Incomplete}}</code>\n已跳过: <code>{{.Skipped}}</code>"
batch_failed_item: "<b>❌ 处理失败</b>\n\n失败文件: <code>{{.Index}}. {{.Name}}</code>\n失败阶段: <code>{{.Stage}}</code>\n失败进度: <code>{{.Progress}}</code>\n失败前速度: <code>{{.Speed}}</code>\n原因: <code>{{.Reason}}</code>\n\n✅ 已完成: <code>{{.Completed}}</code>\n❌ 失败: <code>{{.Failed}}</code>\n⏹ 未完成: <code>{{.Incomplete}}</code>"
batch_failed_group: "<b>❌ 批量上传失败</b>\n\n受影响文件: <code>{{.Affected}} 个</code>\n原因: <code>{{.Reason}}</code>\n\n✅ 已完成: <code>{{.Completed}}</code>\n❌ 批次失败: <code>{{.Failed}}</code>\n⏹ 未完成: <code>{{.Incomplete}}</code>"
batch_failed_task: "<b>❌ 处理失败</b>\n\n原因: <code>{{.Reason}}</code>\n\n✅ 已完成: <code>{{.Completed}}</code>\n⏹ 未完成: <code>{{.Incomplete}}</code>"
batch_failure_stage_download: "下载"
batch_failure_stage_cache: "本地缓存"
batch_failure_stage_upload: "上传"
batch_failure_stage_confirm: "云端确认"
batch_failure_stage_batch_upload: "批量上传"
batch_failure_stage_internal: "任务内部"
single_status_header: "<b>📦 正在处理</b>"
single_downloading: "<blockquote><b>⬇️ 下载中</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / <code>{{.Size}}</code>\n保存至<code>{{.Destination}}</code></blockquote>"
single_downloading_unknown: "<blockquote><b>⬇️ 下载中</b>\n<code>{{.Name}}</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / 未知\n保存至<code>{{.Destination}}</code></blockquote>"
single_uploading: "<blockquote><b>⬆️ 上传中</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / <code>{{.Size}}</code>\n保存至<code>{{.Destination}}</code></blockquote>"
single_upload_retrying: "<blockquote><b>🔁 上传重试</b>\n<code>{{.Name}}</code>\n{{.Bar}} <code>{{.Progress}}%</code>\n尝试次数<code>{{.Attempt}}</code>\n速度<code>{{.Speed}}</code>\n大小<code>{{.Current}}</code> / <code>{{.Size}}</code>\n保存至<code>{{.Destination}}</code></blockquote>"
single_done: "<b>✅ 处理完成</b>\n\n文件名<code>{{.Name}}</code>\n总大小<code>{{.Size}}</code>\n保存至<code>{{.Destination}}</code>"
single_canceled: "<b>🚫 任务已取消</b>\n\n文件名<code>{{.Name}}</code>"
single_failed: "<b>❌ 处理失败</b>\n\n文件名<code>{{.Name}}</code>\n原因<code>{{.Reason}}</code>"
batch_start_prefix: "开始执行批量下载任务\n总大小: "
batch_processing_prefix: "正在处理批量下载任务\n总大小: "
downloading_prefix: "正在下载\n总大小: "
size_with_files: "{{.Size}} ({{.Count}} 个文件)"
size_with_resources: "{{.Size}} ({{.Count}} 个资源)"
processing_list_prefix: "\n正在处理:\n"
processing_none: " - 无"
avg_speed_prefix: "\n平均速度: "
current_progress_prefix: "\n当前进度: "
task_canceled: "任务已取消"
task_canceled_with_id: "处理已取消: {{.TaskID}}"
task_failed_with_error: "处理失败: {{.Error}}"
batch_done_prefix: "处理完成\n文件数: "
direct_done_prefix: "处理完成, 文件数量: "
parsed_start_prefix: "开始下载 {{.Site}} 的资源\n总大小: "
parsed_done_prefix: "处理完成, 资源数量: "
telegraph_start_prefix: "开始下载Telegraph\n图片数量: "
telegraph_progress_prefix: "正在下载\n当前进度: "
telegraph_done_prefix: "处理完成\n图片数量: "
file_start_prefix: "开始下载\n文件名: "
file_processing_prefix: "正在处理下载任务\n文件名: "
download_failed_prefix: "下载失败\n文件名: "
download_done_prefix: "下载完成\n文件名: "
file_size_prefix: "\n文件大小: "
save_path_prefix: "\n保存路径: "
total_size_prefix: "\n总大小: "
direct_start: "开始下载, 总大小: {{.SizeMB}} MB ({{.Count}} 个文件)"
file_name_prefix: "文件名: "
error_prefix: "\n错误: "
aria2_start: "等待 Aria2 下载完成 (GID: {{.GID}})..."
aria2_downloading: "Aria2 正在下载 (GID: {{.GID}})\n"

View File

@@ -1,10 +1,7 @@
package tdler
import (
"context"
"github.com/gotd/td/telegram/downloader"
"github.com/gotd/td/tg"
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/pkg/consts/tglimit"
@@ -13,23 +10,5 @@ import (
func NewDownloader(file tfile.TGFile) *downloader.Builder {
return downloader.NewDownloader().WithPartSize(tglimit.MaxPartSize).
Download(eofAwareClient{Client: file.Dler(), size: file.Size()}, file.Location()).
WithThreads(dlutil.BestThreads(file.Size(), config.C().Threads))
}
// eofAwareClient answers upload.getFile requests at or past the end of the
// file with an empty chunk. gotd's downloader is size-unaware: for files
// whose size is an exact multiple of the part size it issues one final
// request at offset == size and expects an empty chunk, but Telegram rejects
// it with 400 OFFSET_INVALID and the whole download fails.
type eofAwareClient struct {
downloader.Client
size int64
}
func (c eofAwareClient) UploadGetFile(ctx context.Context, req *tg.UploadGetFileRequest) (tg.UploadFileClass, error) {
if req.Offset >= c.size {
return &tg.UploadFile{}, nil
}
return c.Client.UploadGetFile(ctx, req)
Download(file.Dler(), file.Location()).WithThreads(dlutil.BestThreads(file.Size(), config.C().Threads))
}

View File

@@ -1,112 +0,0 @@
package tdler
import (
"bytes"
"context"
"sync"
"testing"
"github.com/gotd/td/tg"
"github.com/gotd/td/tgerr"
"github.com/krau/SaveAny-Bot/pkg/tfile"
)
// serverLikeClient mimics real Telegram upload.getFile behavior: it returns
// up to limit bytes per chunk, and answers any offset at or past the end of
// the file with 400 OFFSET_INVALID.
type serverLikeClient struct {
data []byte
mu sync.Mutex
maxOffset int64
}
func (c *serverLikeClient) UploadGetFile(_ context.Context, req *tg.UploadGetFileRequest) (tg.UploadFileClass, error) {
c.mu.Lock()
if req.Offset > c.maxOffset {
c.maxOffset = req.Offset
}
c.mu.Unlock()
if req.Offset >= int64(len(c.data)) {
return nil, tgerr.New(400, "OFFSET_INVALID")
}
end := min(len(c.data), int(req.Offset)+req.Limit)
return &tg.UploadFile{Bytes: c.data[req.Offset:end]}, nil
}
func (c *serverLikeClient) UploadGetFileHashes(context.Context, *tg.UploadGetFileHashesRequest) ([]tg.FileHash, error) {
return nil, nil
}
func (c *serverLikeClient) UploadReuploadCDNFile(context.Context, *tg.UploadReuploadCDNFileRequest) ([]tg.FileHash, error) {
return nil, nil
}
func (c *serverLikeClient) UploadGetCDNFileHashes(context.Context, *tg.UploadGetCDNFileHashesRequest) ([]tg.FileHash, error) {
return nil, nil
}
func (c *serverLikeClient) UploadGetWebFile(context.Context, *tg.UploadGetWebFileRequest) (*tg.UploadWebFile, error) {
return nil, nil
}
type memWriterAt struct {
b []byte
}
func (w *memWriterAt) WriteAt(p []byte, off int64) (int, error) {
copy(w.b[off:], p)
return len(p), nil
}
func TestDownloadServerLikeEOF(t *testing.T) {
const partSize = 1024 * 1024
tests := []struct {
name string
size int
parallel bool
}{
{"stream exact multiple of part size", 2 * partSize, false},
{"stream non-multiple", 2*partSize + 12345, false},
{"stream smaller than part size", 1234, false},
{"parallel exact multiple of part size", 2 * partSize, true},
{"parallel non-multiple", 2*partSize + 12345, true},
{"parallel smaller than part size", 1234, true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
data := make([]byte, tt.size)
for i := range data {
data[i] = byte(i % 251)
}
client := &serverLikeClient{data: data}
file := tfile.NewTGFile(
&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2},
client, int64(tt.size), "test.bin",
)
dl := NewDownloader(file)
var got []byte
var err error
if tt.parallel {
buf := make([]byte, tt.size)
_, err = dl.WithThreads(4).Parallel(context.Background(), &memWriterAt{b: buf})
got = buf
} else {
var buf bytes.Buffer
_, err = dl.Stream(context.Background(), &buf)
got = buf.Bytes()
}
if err != nil {
t.Fatalf("download failed: %v", err)
}
if !bytes.Equal(got, data) {
t.Fatalf("downloaded %d bytes, want %d matching bytes", len(got), len(data))
}
if client.maxOffset >= int64(tt.size) {
t.Fatalf("requested offset %d at or past EOF (size %d)", client.maxOffset, tt.size)
}
})
}
}

View File

@@ -1,257 +0,0 @@
package tdler
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"os"
"strings"
"sync"
"github.com/gotd/td/tg"
"github.com/gotd/td/tgerr"
"golang.org/x/sync/errgroup"
"github.com/krau/SaveAny-Bot/pkg/consts/tglimit"
"github.com/krau/SaveAny-Bot/pkg/tfile"
)
const maxChunkRetries = 20
// resumeBitmap records which partSize-aligned blocks of a download have been
// durably written, so an interrupted download can continue after a restart.
// The bitmap file is rewritten atomically after every completed block.
type resumeBitmap struct {
PartSize int `json:"part_size"`
Size int64 `json:"size"`
Blocks []uint64 `json:"blocks"`
mu sync.Mutex
}
func newResumeBitmap(size int64) *resumeBitmap {
b := &resumeBitmap{PartSize: tglimit.MaxPartSize, Size: size}
b.ensureBlocks()
return b
}
func (b *resumeBitmap) ensureBlocks() {
if need := (b.blockCount() + 63) / 64; len(b.Blocks) < need {
b.Blocks = make([]uint64, need)
}
}
func (b *resumeBitmap) blockCount() int {
return int((b.Size + int64(b.PartSize) - 1) / int64(b.PartSize))
}
func (b *resumeBitmap) isDone(block int) bool {
return b.Blocks[block/64]&(1<<uint(block%64)) != 0
}
func (b *resumeBitmap) markDone(block int) {
b.Blocks[block/64] |= 1 << uint(block%64)
}
func (b *resumeBitmap) complete() bool {
for block := 0; block < b.blockCount(); block++ {
if !b.isDone(block) {
return false
}
}
return true
}
func (b *resumeBitmap) missingBlocks() []int {
missing := make([]int, 0, b.blockCount())
for block := 0; block < b.blockCount(); block++ {
if !b.isDone(block) {
missing = append(missing, block)
}
}
return missing
}
func loadResumeBitmap(path string) (*resumeBitmap, error) {
data, err := os.ReadFile(path)
if errors.Is(err, os.ErrNotExist) {
return nil, nil
}
if err != nil {
return nil, fmt.Errorf("read resume bitmap: %w", err)
}
var b resumeBitmap
if err := json.Unmarshal(data, &b); err != nil {
// 无法解析的位图 (外部损坏): 删除并视为不存在, 全量重下自愈。
_ = os.Remove(path)
return nil, nil
}
if b.Size <= 0 || b.PartSize <= 0 {
// 无效位图 (损坏或旧格式), 视为不存在, 全量重下。
_ = os.Remove(path)
return nil, nil
}
b.ensureBlocks()
return &b, nil
}
func (b *resumeBitmap) save(path string) error {
b.mu.Lock()
defer b.mu.Unlock()
return b.saveLocked(path)
}
func (b *resumeBitmap) saveLocked(path string) error {
data, err := json.Marshal(b)
if err != nil {
return fmt.Errorf("marshal resume bitmap: %w", err)
}
tmp := path + ".tmp"
if err := os.WriteFile(tmp, data, 0o644); err != nil {
return fmt.Errorf("write resume bitmap: %w", err)
}
return os.Rename(tmp, path)
}
func (b *resumeBitmap) markAndSave(block int, path string) error {
b.mu.Lock()
defer b.mu.Unlock()
b.markDone(block)
return b.saveLocked(path)
}
func isRetryableTimeout(ctx context.Context, err error) bool {
if err == nil || ctx.Err() != nil {
return false
}
if tgerr.Is(err, tg.ErrTimeout) || errors.Is(err, context.DeadlineExceeded) {
return true
}
var netErr net.Error
return errors.As(err, &netErr) && netErr.Timeout()
}
// fetchChunk downloads one partSize-aligned chunk, retrying flood waits and
// transient timeouts like gotd's downloader does.
func fetchChunk(ctx context.Context, file tfile.TGFile, offset int64, limit int) ([]byte, error) {
req := &tg.UploadGetFileRequest{
Location: file.Location(),
Offset: offset,
Limit: limit,
}
timeoutRetries := 0
for {
res, err := file.Dler().UploadGetFile(ctx, req)
if err == nil {
switch r := res.(type) {
case *tg.UploadFile:
return r.Bytes, nil
case *tg.UploadFileCDNRedirect:
return nil, fmt.Errorf("CDN redirect is not supported (dc %d)", r.DCID)
default:
return nil, fmt.Errorf("unexpected upload.getFile response %T", res)
}
}
if flood, ferr := tgerr.FloodWait(ctx, err); ferr != nil {
if flood {
// FloodWait already slept; retry.
continue
}
if isRetryableTimeout(ctx, ferr) {
timeoutRetries++
if timeoutRetries >= maxChunkRetries {
return nil, fmt.Errorf("get chunk at %d: retry limit reached: %w", offset, ferr)
}
continue
}
return nil, fmt.Errorf("get chunk at %d: %w", offset, ferr)
}
}
}
// DownloadResumable downloads file to w in partSize chunks, skipping blocks
// already recorded as complete in bitmapPath and persisting every completed
// block so an interrupted download can resume. A missing or incompatible
// bitmap starts a full download. Requires a known, non-zero file size.
func DownloadResumable(
ctx context.Context,
file tfile.TGFile,
w io.WriterAt,
threads int,
bitmapPath string,
) error {
if file.Size() <= 0 {
return fmt.Errorf("resumable download requires a known size")
}
bm, err := loadResumeBitmap(bitmapPath)
if err != nil {
return err
}
// 位图描述的数据文件 (bitmapPath 去掉 .bitmap 后缀) 必须存在且非空:
// 若缺失或为空, 已标记完成的块字节已丢失, 必须重置位图全量重下。
if bm != nil {
partPath := strings.TrimSuffix(bitmapPath, ".bitmap")
if stat, err := os.Stat(partPath); err != nil || stat.Size() == 0 {
if err := os.Remove(bitmapPath); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("reset stale resume bitmap: %w", err)
}
bm = nil
}
}
if bm == nil || bm.PartSize != tglimit.MaxPartSize || bm.Size != file.Size() {
bm = newResumeBitmap(file.Size())
if err := bm.save(bitmapPath); err != nil {
return err
}
}
missing := bm.missingBlocks()
if len(missing) == 0 {
return nil
}
eg, gctx := errgroup.WithContext(ctx)
eg.SetLimit(threads)
for _, block := range missing {
block := block
eg.Go(func() error {
offset := int64(block) * int64(bm.PartSize)
data, err := fetchChunk(gctx, file, offset, bm.PartSize)
if err != nil {
return err
}
if len(data) == 0 {
return fmt.Errorf("file ended early at offset %d (expected size %d)", offset, bm.Size)
}
if _, err := w.WriteAt(data, offset); err != nil {
return fmt.Errorf("write chunk at offset %d: %w", offset, err)
}
return bm.markAndSave(block, bitmapPath)
})
}
if err := eg.Wait(); err != nil {
return err
}
if !bm.complete() {
return fmt.Errorf("download finished with missing blocks")
}
return nil
}
// RemoveResumeState deletes the bitmap file of a completed download.
func RemoveResumeState(bitmapPath string) error {
if err := os.Remove(bitmapPath); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("remove resume bitmap: %w", err)
}
if err := os.Remove(bitmapPath + ".tmp"); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("remove resume bitmap temp: %w", err)
}
return nil
}
// ResumeStatePath returns the bitmap path for a download cache file.
func ResumeStatePath(cachePath string) string {
return cachePath + ".bitmap"
}

View File

@@ -1,271 +0,0 @@
package tdler
import (
"context"
"errors"
"os"
"path/filepath"
"testing"
"github.com/gotd/td/tg"
"github.com/gotd/td/tgerr"
"github.com/krau/SaveAny-Bot/pkg/tfile"
)
// failAfterClient serves the first failAfter chunks, then returns err.
type failAfterClient struct {
*serverLikeClient
failAfter int
calls int
err error
}
func (c *failAfterClient) UploadGetFile(ctx context.Context, req *tg.UploadGetFileRequest) (tg.UploadFileClass, error) {
c.calls++
if c.calls > c.failAfter {
return nil, c.err
}
return c.serverLikeClient.UploadGetFile(ctx, req)
}
func TestDownloadResumableFull(t *testing.T) {
data := make([]byte, 3*1024*1024+123)
for i := range data {
data[i] = byte(i % 251)
}
client := &serverLikeClient{data: data}
file := tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, client, int64(len(data)), "test.bin")
dir := t.TempDir()
bitmapPath := filepath.Join(dir, "test.bin.bitmap")
w := &memWriterAt{b: make([]byte, len(data))}
if err := DownloadResumable(context.Background(), file, w, 4, bitmapPath); err != nil {
t.Fatalf("download failed: %v", err)
}
if !bytesEqual(w.b, data) {
t.Fatalf("downloaded data mismatch")
}
bm, err := loadResumeBitmap(bitmapPath)
if err != nil {
t.Fatalf("load bitmap: %v", err)
}
if bm == nil || !bm.complete() {
t.Fatalf("bitmap not complete after full download")
}
}
func TestDownloadResumableInterrupted(t *testing.T) {
data := make([]byte, 5*1024*1024) // exactly 5 blocks
for i := range data {
data[i] = byte(i % 251)
}
dir := t.TempDir()
bitmapPath := filepath.Join(dir, "test.bin.bitmap")
w := &memWriterAt{b: make([]byte, len(data))}
// First run: 3 blocks complete, 4th request fails.
flaky := &failAfterClient{
serverLikeClient: &serverLikeClient{data: data},
failAfter: 3,
err: tgerr.New(500, "INTERNAL_SERVER_ERROR"),
}
file := tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, flaky, int64(len(data)), "test.bin")
err := DownloadResumable(context.Background(), file, w, 1, bitmapPath)
if err == nil {
t.Fatalf("expected first run to fail")
}
bm, err := loadResumeBitmap(bitmapPath)
if err != nil {
t.Fatalf("load bitmap after interruption: %v", err)
}
if bm == nil {
t.Fatalf("bitmap missing after interruption")
}
if got := bm.blockCount() - len(bm.missingBlocks()); got != 3 {
t.Fatalf("expected 3 completed blocks, got %d", got)
}
// Second run: only the missing blocks are requested.
healthy := &serverLikeClient{data: data}
file = tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, healthy, int64(len(data)), "test.bin")
if err := DownloadResumable(context.Background(), file, w, 1, bitmapPath); err != nil {
t.Fatalf("resume failed: %v", err)
}
if !bytesEqual(w.b, data) {
t.Fatalf("resumed data mismatch")
}
if healthy.maxOffset >= int64(len(data)) {
t.Fatalf("resume requested offset %d at or past EOF", healthy.maxOffset)
}
if bm, err = loadResumeBitmap(bitmapPath); err != nil || bm == nil || !bm.complete() {
t.Fatalf("bitmap not complete after resume: %v", err)
}
}
func TestDownloadResumableBitmapResetOnSizeChange(t *testing.T) {
data := make([]byte, 2*1024*1024)
for i := range data {
data[i] = byte(i % 251)
}
dir := t.TempDir()
bitmapPath := filepath.Join(dir, "test.bin.bitmap")
w := &memWriterAt{b: make([]byte, len(data))}
// Record a bitmap claiming the old, larger file is fully downloaded.
stale := newResumeBitmap(int64(4 * 1024 * 1024))
if err := stale.save(bitmapPath); err != nil {
t.Fatalf("save stale bitmap: %v", err)
}
client := &serverLikeClient{data: data}
file := tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, client, int64(len(data)), "test.bin")
if err := DownloadResumable(context.Background(), file, w, 1, bitmapPath); err != nil {
t.Fatalf("download with stale bitmap failed: %v", err)
}
if !bytesEqual(w.b, data) {
t.Fatalf("data mismatch with stale bitmap")
}
}
// TestDownloadResumablePartMissingOrTruncated resets the bitmap: skipped
// blocks would otherwise be zero-filled (caller recreates the part file
// without its bytes), or the download would wedge forever on a stale
// complete bitmap.
func TestDownloadResumablePartMissingOrTruncated(t *testing.T) {
data := make([]byte, 5*1024*1024)
for i := range data {
data[i] = byte(i % 251)
}
dir := t.TempDir()
partPath := filepath.Join(dir, "test.bin.part")
bitmapPath := ResumeStatePath(partPath)
tests := []struct {
name string
doneBlocks []int
createPart bool
truncate bool
}{
{"part missing, partial bitmap", []int{0, 1, 2}, false, false},
{"part empty, partial bitmap", []int{0, 1, 2}, true, true},
{"part missing, complete bitmap", []int{0, 1, 2, 3, 4}, false, false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
os.Remove(partPath)
os.Remove(bitmapPath)
bm := newResumeBitmap(int64(len(data)))
for _, block := range tt.doneBlocks {
bm.markDone(block)
}
if err := bm.save(bitmapPath); err != nil {
t.Fatal(err)
}
if tt.createPart {
// Simulate the caller re-creating the part file (truncating).
if err := os.WriteFile(partPath, nil, 0o644); err != nil {
t.Fatal(err)
}
if tt.truncate {
if err := os.WriteFile(partPath, make([]byte, 0), 0o644); err != nil {
t.Fatal(err)
}
}
}
partFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644)
if err != nil {
t.Fatal(err)
}
defer partFile.Close()
client := &serverLikeClient{data: data}
file := tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, client, int64(len(data)), "test.bin")
if err := DownloadResumable(context.Background(), file, partFile, 1, bitmapPath); err != nil {
t.Fatalf("download failed: %v", err)
}
got := make([]byte, len(data))
if _, err := partFile.ReadAt(got, 0); err != nil {
t.Fatal(err)
}
if !bytesEqual(got, data) {
t.Fatalf("downloaded data mismatch (blocks not reset)")
}
})
}
}
// TestDownloadResumableInvalidBitmap treats a corrupt bitmap as absent.
func TestDownloadResumableInvalidBitmap(t *testing.T) {
data := make([]byte, 1024*1024+7)
for i := range data {
data[i] = byte(i % 251)
}
dir := t.TempDir()
partPath := filepath.Join(dir, "test.bin.part")
bitmapPath := ResumeStatePath(partPath)
for _, content := range []string{
`{"part_size":1048576,"size":-1,"blocks":[]}`,
`{"part_size":1048576,"size":9223372036854775807,"blocks":[]}`,
`not json`,
} {
os.Remove(partPath)
if err := os.WriteFile(bitmapPath, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
partFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644)
if err != nil {
t.Fatal(err)
}
client := &serverLikeClient{data: data}
file := tfile.NewTGFile(&tg.InputDocumentFileLocation{ID: 1, AccessHash: 2}, client, int64(len(data)), "test.bin")
err = DownloadResumable(context.Background(), file, partFile, 1, bitmapPath)
partFile.Close()
if err != nil {
t.Fatalf("download with corrupt bitmap %q failed: %v", content, err)
}
got := make([]byte, len(data))
f, err := os.Open(partPath)
if err != nil {
t.Fatal(err)
}
if _, err := f.ReadAt(got, 0); err != nil {
t.Fatal(err)
}
f.Close()
if !bytesEqual(got, data) {
t.Fatalf("downloaded data mismatch with corrupt bitmap %q", content)
}
}
}
func TestRemoveResumeState(t *testing.T) {
dir := t.TempDir()
path := filepath.Join(dir, "x.bitmap")
if err := os.WriteFile(path, []byte("{}"), 0o644); err != nil {
t.Fatal(err)
}
if err := RemoveResumeState(path); err != nil {
t.Fatalf("RemoveResumeState: %v", err)
}
if _, err := os.Stat(path); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("bitmap still exists: %v", err)
}
// Removing again must be a no-op.
if err := RemoveResumeState(path); err != nil {
t.Fatalf("RemoveResumeState second call: %v", err)
}
}
func bytesEqual(a, b []byte) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}

View File

@@ -1,42 +0,0 @@
package fsutil_test
import (
"errors"
"os"
"path/filepath"
"testing"
"github.com/krau/SaveAny-Bot/common/utils/fsutil"
)
func TestCloseAndRemove(t *testing.T) {
tests := []struct {
name string
preClose bool
}{
{name: "open file"},
{name: "already closed file", preClose: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
filePath := filepath.Join(t.TempDir(), "cache-file")
file, err := fsutil.CreateFile(filePath)
if err != nil {
t.Fatalf("CreateFile() failed: %v", err)
}
if tt.preClose {
if err := file.Close(); err != nil {
t.Fatalf("Close() failed: %v", err)
}
}
if err := file.CloseAndRemove(); err != nil {
t.Fatalf("CloseAndRemove() failed: %v", err)
}
if _, err := os.Stat(filePath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("cache file still exists after CloseAndRemove(): %v", err)
}
})
}
}

View File

@@ -1,7 +1,6 @@
package fsutil
import (
"errors"
"os"
"path/filepath"
"strings"
@@ -42,11 +41,10 @@ func (f *File) Remove() error {
}
func (f *File) CloseAndRemove() error {
closeErr := f.Close()
if errors.Is(closeErr, os.ErrClosed) {
closeErr = nil
if err := f.Close(); err != nil {
return err
}
return errors.Join(closeErr, f.Remove())
return f.Remove()
}
func CreateFile(fp string) (*File, error) {

View File

@@ -1,27 +0,0 @@
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))
}

View File

@@ -17,11 +17,7 @@ type ProgressReadSeeker struct {
// Seek implements io.ReadSeeker.
func (pr *ProgressReadSeeker) Seek(offset int64, whence int) (int64, error) {
position, err := pr.reader.Seek(offset, whence)
if err == nil {
pr.read.Store(position)
}
return position, err
return pr.reader.Seek(offset, whence)
}
// NewProgressReader creates a new ProgressReader
@@ -58,7 +54,7 @@ func (pr *ProgressReadSeeker) Progress() float64 {
return float64(pr.read.Load()) / float64(pr.total.Load())
}
// BytesRead returns the current tracked reader position.
// Read returns the number of bytes read so far
func (pr *ProgressReadSeeker) BytesRead() int64 {
return pr.read.Load()
}

View File

@@ -1,50 +0,0 @@
package ioutil
import (
"bytes"
"io"
"testing"
)
func TestProgressReadSeekerTracksReads(t *testing.T) {
var gotRead, gotTotal int64
reader := NewProgressReader(bytes.NewReader([]byte("abcdef")), 6, func(read, total int64) {
gotRead = read
gotTotal = total
})
buffer := make([]byte, 4)
if _, err := io.ReadFull(reader, buffer); err != nil {
t.Fatalf("read failed: %v", err)
}
if gotRead != 4 || gotTotal != 6 {
t.Fatalf("progress = %d/%d, want 4/6", gotRead, gotTotal)
}
if reader.BytesRead() != 4 {
t.Fatalf("BytesRead() = %d, want 4", reader.BytesRead())
}
}
func TestProgressReadSeekerResetsPositionOnSeek(t *testing.T) {
reader := NewProgressReader(bytes.NewReader([]byte("abcdef")), 6, nil)
buffer := make([]byte, 4)
if _, err := io.ReadFull(reader, buffer); err != nil {
t.Fatalf("read failed: %v", err)
}
position, err := reader.Seek(0, io.SeekStart)
if err != nil {
t.Fatalf("seek failed: %v", err)
}
if position != 0 || reader.BytesRead() != 0 {
t.Fatalf("position after seek = %d (tracked %d), want 0", position, reader.BytesRead())
}
if _, err := io.ReadFull(reader, buffer[:2]); err != nil {
t.Fatalf("read after seek failed: %v", err)
}
if reader.BytesRead() != 2 {
t.Fatalf("BytesRead() after seek and read = %d, want 2", reader.BytesRead())
}
}

View File

@@ -1,50 +0,0 @@
// 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
}

View File

@@ -1,35 +0,0 @@
package tgutil
import (
stdhtml "html"
"github.com/gotd/td/telegram/message/entity"
messagehtml "github.com/gotd/td/telegram/message/html"
"github.com/gotd/td/telegram/message/styling"
"github.com/gotd/td/tg"
)
// EscapeHTMLTemplateData returns a copy of data with string values escaped for
// interpolation into Telegram HTML templates.
func EscapeHTMLTemplateData(data map[string]any) map[string]any {
escaped := make(map[string]any, len(data))
for key, value := range data {
if text, ok := value.(string); ok {
escaped[key] = stdhtml.EscapeString(text)
continue
}
escaped[key] = value
}
return escaped
}
// RenderHTML renders Telegram-compatible HTML into plain text and message
// entities.
func RenderHTML(markup string) (string, []tg.MessageEntityClass, error) {
var builder entity.Builder
if err := styling.Perform(&builder, messagehtml.String(nil, markup)); err != nil {
return "", nil, err
}
text, entities := builder.Complete()
return text, entities, nil
}

View File

@@ -1,54 +0,0 @@
package tgutil
import (
"testing"
"github.com/gotd/td/tg"
)
func TestEscapeHTMLTemplateDataDoesNotMutateInput(t *testing.T) {
input := map[string]any{
"Text": `<b>A&B</b>`,
"Count": 2,
}
escaped := EscapeHTMLTemplateData(input)
if got, want := escaped["Text"], "&lt;b&gt;A&amp;B&lt;/b&gt;"; got != want {
t.Fatalf("escaped text = %q, want %q", got, want)
}
if got := input["Text"]; got != `<b>A&B</b>` {
t.Fatalf("input was mutated: %q", got)
}
if got := escaped["Count"]; got != 2 {
t.Fatalf("non-string value = %v, want 2", got)
}
}
func TestRenderHTMLUsesTemplateStylesAndDecodesValues(t *testing.T) {
data := EscapeHTMLTemplateData(map[string]any{"Name": `<b>A&B</b>.bin`})
markup := `<blockquote><b>Uploading</b>
<code>` + data["Name"].(string) + `</code></blockquote>`
text, entities, err := RenderHTML(markup)
if err != nil {
t.Fatalf("RenderHTML() failed: %v", err)
}
if want := "Uploading\n<b>A&B</b>.bin"; text != want {
t.Fatalf("rendered text = %q, want %q", text, want)
}
var bold, code, blockquote int
for _, messageEntity := range entities {
switch messageEntity.(type) {
case *tg.MessageEntityBold:
bold++
case *tg.MessageEntityCode:
code++
case *tg.MessageEntityBlockquote:
blockquote++
}
}
if bold != 1 || code != 1 || blockquote != 1 {
t.Fatalf("entity counts = bold:%d code:%d blockquote:%d", bold, code, blockquote)
}
}

View File

@@ -2,7 +2,6 @@ package tgutil
import (
"fmt"
"sort"
"strconv"
"strings"
"unicode"
@@ -194,6 +193,97 @@ 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 {
@@ -269,16 +359,9 @@ func GetGroupedMessages(ctx *ext.Context, chatID int64, msg *tg.Message) ([]*tg.
groupedMessages = append(groupedMessages, m)
}
}
sortMessagesByID(groupedMessages)
return groupedMessages, nil
}
func sortMessagesByID(messages []*tg.Message) {
sort.Slice(messages, func(i, j int) bool {
return messages[i].GetID() < messages[j].GetID()
})
}
func ExtractMessageEntityUrls(msg *tg.Message) []string {
if len(msg.Entities) == 0 {
return nil

View File

@@ -1,18 +0,0 @@
package tgutil
import (
"testing"
"github.com/gotd/td/tg"
)
func TestSortMessagesByID(t *testing.T) {
messages := []*tg.Message{{ID: 9}, {ID: 3}, {ID: 7}}
sortMessagesByID(messages)
want := []int{3, 7, 9}
for i := range messages {
if messages[i].GetID() != want[i] {
t.Fatalf("message %d has ID %d, want %d", i, messages[i].GetID(), want[i])
}
}
}

View File

@@ -44,25 +44,6 @@ 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

View File

@@ -10,4 +10,13 @@ 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"`
// }

View File

@@ -8,16 +8,14 @@ import (
type TelegramStorageConfig struct {
BaseConfig
ChatID int64 `toml:"chat_id" mapstructure:"chat_id" json:"chat_id"`
ForceFile bool `toml:"force_file" mapstructure:"force_file" json:"force_file"`
RateLimit int `toml:"rate_limit" mapstructure:"rate_limit" json:"rate_limit"`
RateBurst int `toml:"rate_burst" mapstructure:"rate_burst" json:"rate_burst"`
SkipLarge bool `toml:"skip_large" mapstructure:"skip_large" json:"skip_large"` // skip files larger than Telegram limit(2GB)
SplitLargeVideo bool `toml:"split_large_video" mapstructure:"split_large_video" json:"split_large_video"`
// split files larger than the uploader account limit into parts of specified size, in MB
// leave 0 to use the account limit (2000MB for bots/regular users, 4000MB for Premium users)
ChatID int64 `toml:"chat_id" mapstructure:"chat_id" json:"chat_id"`
ForceFile bool `toml:"force_file" mapstructure:"force_file" json:"force_file"`
RateLimit int `toml:"rate_limit" mapstructure:"rate_limit" json:"rate_limit"`
RateBurst int `toml:"rate_burst" mapstructure:"rate_burst" json:"rate_burst"`
SkipLarge bool `toml:"skip_large" mapstructure:"skip_large" json:"skip_large"` // skip files larger than Telegram limit(2GB)
// split files larger than Telegram limit(2GB) into parts of specified size, in MB, leave 0 to set default(2000MB)
// only effective when SkipLarge is false
// use zip when splitting non-video files or when lossless video splitting is disabled/unavailable
// use zip when splitting
SplitSizeMB int64 `toml:"split_size_mb" mapstructure:"split_size_mb" json:"split_size_mb"`
}

View File

@@ -9,7 +9,6 @@ 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"
@@ -69,13 +68,6 @@ 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()
@@ -84,13 +76,11 @@ 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://") {
// 使用远程配置文件
client := &http.Client{Timeout: 30 * time.Second}
resp, err := client.Get(cfg)
resp, err := http.Get(cfg)
if err != nil {
return fmt.Errorf("failed to fetch remote config file: %w", err)
}
@@ -101,7 +91,6 @@ 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)
}
@@ -152,15 +141,13 @@ func Init(ctx context.Context, configFile ...string) error {
viper.SetDefault(key, value)
}
if !loadedFromURL {
if err := viper.ReadInConfig(); err != nil {
logger.Errorf("Error reading config file: %v", err)
return err
}
if err := viper.ReadInConfig(); err != nil {
fmt.Println("Error reading config file, ", err)
return err
}
if err := viper.Unmarshal(cfg); err != nil {
logger.Errorf("Error unmarshalling config file: %v", err)
fmt.Println("Error unmarshalling config file, ", err)
return err
}

View File

@@ -3,28 +3,15 @@ package core
import (
"context"
"errors"
"sync"
"github.com/charmbracelet/log"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/database"
"github.com/krau/SaveAny-Bot/pkg/enums/tasktype"
"github.com/krau/SaveAny-Bot/pkg/queue"
"github.com/krau/SaveAny-Bot/pkg/taskevent"
)
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
}
var queueInstance *queue.TaskQueue[Executable]
type Executable interface {
Type() tasktype.TaskType
@@ -46,9 +33,6 @@ func worker(ctx context.Context, qe *queue.TaskQueue[Executable], semaphore chan
exe := qtask.Data
taskCtx := qtask.Context()
logger.Infof("Processing task: %s", exe.TaskID())
if err := database.UpdateTaskStatus(taskCtx, exe.TaskID(), database.TaskStatusRunning, ""); err != nil {
logger.Errorf("Failed to mark task %s as running: %v", exe.TaskID(), err)
}
taskevent.Emit(taskCtx, taskevent.Event{TaskID: exe.TaskID(), Phase: taskevent.PhaseStart})
if err := ExecCommandString(taskCtx, execHooks.TaskBeforeStart); err != nil {
logger.Errorf("Failed to execute before start hook for task %s: %v", exe.TaskID(), err)
@@ -74,11 +58,6 @@ func worker(ctx context.Context, qe *queue.TaskQueue[Executable], semaphore chan
}
taskevent.Emit(taskCtx, taskevent.Event{TaskID: exe.TaskID(), Phase: taskevent.PhaseDone, Err: err})
qe.Done(qtask.ID)
// 用独立 ctx 删除: 优雅关停时 run ctx 已被取消, 会留下已完成任务的行,
// 导致重启后重复执行 (重复上传)。
if err := database.DeleteTask(context.Background(), exe.TaskID()); err != nil {
logger.Errorf("Failed to delete persisted task %s: %v", exe.TaskID(), err)
}
<-semaphore
}
}
@@ -86,36 +65,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)
q := initQueue()
if queueInstance == nil {
queueInstance = queue.NewTaskQueue[Executable]()
}
for range config.C().Workers {
go worker(ctx, q, semaphore)
go worker(ctx, queueInstance, 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 {
if err := persistTask(ctx, task); err != nil {
log.FromContext(ctx).Errorf("Failed to persist task %s: %v", task.TaskID(), err)
}
return initQueue().Add(queue.NewTask(ctx, task.TaskID(), task.Title(), task))
return queueInstance.Add(queue.NewTask(ctx, task.TaskID(), task.Title(), task))
}
func CancelTask(ctx context.Context, id string) error {
err := queueInstance.CancelTask(id)
if err != nil {
return err
}
if err := database.DeleteTask(ctx, id); err != nil {
log.FromContext(ctx).Errorf("Failed to delete persisted task %s: %v", id, err)
}
return nil
return err
}
func GetLength(ctx context.Context) int {

View File

@@ -1,145 +0,0 @@
package core
import (
"context"
"sync"
"time"
"fmt"
"github.com/charmbracelet/log"
"github.com/gotd/td/telegram/downloader"
"github.com/krau/SaveAny-Bot/database"
"github.com/krau/SaveAny-Bot/pkg/enums/tasktype"
)
// TaskCodec serializes and rebuilds a task from its persisted payload.
// Task types without a registered codec are dropped with a warning on
// recovery instead of being silently re-enqueued.
type TaskCodec interface {
Marshal(task Executable) ([]byte, error)
Unmarshal(payload []byte) (Executable, error)
}
var (
taskCodecsMu sync.RWMutex
taskCodecs = make(map[tasktype.TaskType]TaskCodec)
dlerMu sync.RWMutex
dlerProvider func() downloader.Client
)
func RegisterTaskCodec(t tasktype.TaskType, codec TaskCodec) {
taskCodecsMu.Lock()
defer taskCodecsMu.Unlock()
taskCodecs[t] = codec
}
func TaskCodecFor(t tasktype.TaskType) (TaskCodec, bool) {
taskCodecsMu.RLock()
defer taskCodecsMu.RUnlock()
codec, ok := taskCodecs[t]
return codec, ok
}
// SetDownloaderProvider registers the download client factory used to
// rebuild tfile.TGFile values when recovering tasks.
func SetDownloaderProvider(f func() downloader.Client) {
dlerMu.Lock()
defer dlerMu.Unlock()
dlerProvider = f
}
// DownloaderClient returns the registered download client, or nil.
func DownloaderClient() downloader.Client {
dlerMu.RLock()
defer dlerMu.RUnlock()
if dlerProvider == nil {
return nil
}
return dlerProvider()
}
func persistTask(ctx context.Context, task Executable) error {
codec, ok := TaskCodecFor(task.Type())
if !ok {
return nil
}
payload, err := codec.Marshal(task)
if err != nil {
return err
}
return database.UpsertTask(ctx, &database.Task{
ID: task.TaskID(),
Type: string(task.Type()),
Payload: payload,
Status: string(database.TaskStatusQueued),
})
}
// UpdateTaskPayload atomically mutates the persisted payload of a running
// task (e.g. recording per-element upload progress for recovery).
func UpdateTaskPayload(ctx context.Context, id string, mutate func(payload []byte) ([]byte, error)) error {
row, err := database.GetTask(ctx, id)
if err != nil {
return err
}
updated, err := mutate(row.Payload)
if err != nil {
return fmt.Errorf("mutate payload: %w", err)
}
return database.UpdateTaskPayload(ctx, id, updated)
}
// RecoverTasks re-enqueues tasks that were unfinished when the process last
// exited. Must be called after storages are loaded and before Run. Tasks
// that cannot be recovered are marked failed and kept for visibility.
func RecoverTasks(ctx context.Context) {
logger := log.FromContext(ctx)
if err := database.DeleteStaleFailedTasks(ctx, 24*time.Hour); err != nil {
logger.Warnf("Failed to clean stale failed tasks: %v", err)
}
tasks, err := database.GetUnfinishedTasks(ctx)
if err != nil {
logger.Errorf("Failed to load unfinished tasks: %v", err)
return
}
for _, t := range tasks {
codec, ok := TaskCodecFor(tasktype.TaskType(t.Type))
if !ok {
logger.Warnf("Task %s (type %s) cannot be recovered: no codec registered", t.ID, t.Type)
markRecoverFailed(ctx, t, "no codec registered")
continue
}
task, err := codec.Unmarshal(t.Payload)
if err != nil {
logger.Errorf("Task %s cannot be recovered: failed to rebuild: %v", t.ID, err)
markRecoverFailed(ctx, t, err.Error())
continue
}
if initQueue().Contains(task.TaskID()) {
// Already live in the queue (e.g. submitted via API during
// startup); keep the row as-is.
logger.Infof("Task %s already queued, keeping row", t.ID)
continue
}
if err := AddTask(ctx, task); err != nil {
logger.Errorf("Task %s cannot be recovered: failed to re-enqueue: %v", t.ID, err)
markRecoverFailed(ctx, t, err.Error())
continue
}
// Upsert cleared the original creation time; restore it so
// GetUnfinishedTasks ordering stays stable across restarts.
if err := database.RestoreTaskCreatedAt(ctx, t.ID, t.CreatedAt); err != nil {
logger.Warnf("Failed to restore created_at for task %s: %v", t.ID, err)
}
logger.Infof("Recovered task %s (%s)", t.ID, t.Type)
}
}
func markRecoverFailed(ctx context.Context, t database.Task, reason string) {
if err := database.UpdateTaskStatus(ctx, t.ID, database.TaskStatusFailed, reason); err != nil {
log.FromContext(ctx).Errorf("Failed to mark task %s as failed: %v", t.ID, err)
}
}

View File

@@ -1,162 +0,0 @@
package core
import (
"context"
"fmt"
"os"
"path/filepath"
"testing"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/database"
"github.com/krau/SaveAny-Bot/pkg/enums/tasktype"
)
const testRecoverType = tasktype.TaskType("test-recover")
type stubTask struct {
id string
}
func (s *stubTask) Type() tasktype.TaskType { return testRecoverType }
func (s *stubTask) Title() string { return s.id }
func (s *stubTask) TaskID() string { return s.id }
func (s *stubTask) Execute(context.Context) error { return nil }
type stubCodec struct{}
func (stubCodec) Marshal(task Executable) ([]byte, error) {
return []byte(task.TaskID()), nil
}
func (stubCodec) Unmarshal(payload []byte) (Executable, error) {
if len(payload) == 0 {
return nil, fmt.Errorf("empty payload")
}
return &stubTask{id: string(payload)}, nil
}
func initRecoveryEnv(t *testing.T) context.Context {
t.Helper()
dir := t.TempDir()
cfgPath := filepath.Join(dir, "config.toml")
content := fmt.Sprintf("[db]\npath = %q\n", filepath.Join(dir, "test.db"))
if err := os.WriteFile(cfgPath, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
if err := config.Init(context.Background(), cfgPath); err != nil {
t.Fatalf("config init: %v", err)
}
database.Init(context.Background())
RegisterTaskCodec(testRecoverType, stubCodec{})
return context.Background()
}
func TestRecoverTasksReenqueuesAndMarksUnknownFailed(t *testing.T) {
ctx := initRecoveryEnv(t)
if err := database.CreateTask(ctx, &database.Task{
ID: "rec-1", Type: string(testRecoverType), Payload: []byte("rec-1"), Status: string(database.TaskStatusQueued),
}); err != nil {
t.Fatal(err)
}
if err := database.CreateTask(ctx, &database.Task{
ID: "rec-2", Type: string(testRecoverType), Payload: []byte("rec-2"), Status: string(database.TaskStatusRunning),
}); err != nil {
t.Fatal(err)
}
if err := database.CreateTask(ctx, &database.Task{
ID: "drop-1", Type: "unregistered", Payload: nil, Status: string(database.TaskStatusQueued),
}); err != nil {
t.Fatal(err)
}
RecoverTasks(ctx)
ids := map[string]bool{}
for _, info := range GetQueuedTasks(ctx) {
ids[info.ID] = true
}
if !ids["rec-1"] || !ids["rec-2"] {
t.Fatalf("recovered task ids = %v, want rec-1 and rec-2", ids)
}
unfinished, err := database.GetUnfinishedTasks(ctx)
if err != nil {
t.Fatal(err)
}
if len(unfinished) != 2 {
t.Fatalf("unfinished rows = %d, want 2", len(unfinished))
}
for _, task := range unfinished {
if task.ID == "drop-1" {
t.Fatalf("unregistered task record was not dropped")
}
if task.Status != string(database.TaskStatusQueued) {
t.Fatalf("recovered task status = %s, want queued", task.Status)
}
}
// The unrecoverable task must be kept and marked failed, not silently deleted.
drop, err := database.GetTask(ctx, "drop-1")
if err != nil {
t.Fatalf("dropped task row missing: %v", err)
}
if drop.Status != string(database.TaskStatusFailed) {
t.Fatalf("dropped task status = %s, want failed", drop.Status)
}
if drop.Error == "" {
t.Fatalf("dropped task has no failure reason")
}
}
func TestRecoverTasksMarksInvalidPayloadFailed(t *testing.T) {
ctx := initRecoveryEnv(t)
if err := database.CreateTask(ctx, &database.Task{
ID: "bad-1", Type: string(testRecoverType), Payload: nil, Status: string(database.TaskStatusQueued),
}); err != nil {
t.Fatal(err)
}
RecoverTasks(ctx)
// bad-1 must not be enqueued; its row is kept as failed.
for _, info := range GetQueuedTasks(ctx) {
if info.ID == "bad-1" {
t.Fatalf("task with invalid payload was enqueued")
}
}
count, err := database.CountUnfinishedTasks(ctx)
if err != nil {
t.Fatal(err)
}
if count != 0 {
t.Fatalf("unfinished rows = %d, want 0", count)
}
bad, err := database.GetTask(ctx, "bad-1")
if err != nil {
t.Fatalf("failed task row missing: %v", err)
}
if bad.Status != string(database.TaskStatusFailed) {
t.Fatalf("bad task status = %s, want failed", bad.Status)
}
}
func TestRecoverTasksSkipsAlreadyQueued(t *testing.T) {
ctx := initRecoveryEnv(t)
// A task submitted during startup is both persisted and in the queue.
task := &stubTask{id: "live-1"}
if err := AddTask(ctx, task); err != nil {
t.Fatal(err)
}
RecoverTasks(ctx)
// The row must survive with its original status.
row, err := database.GetTask(ctx, "live-1")
if err != nil {
t.Fatalf("row missing for queued task: %v", err)
}
if row.Status != string(database.TaskStatusQueued) {
t.Fatalf("row status = %s, want queued", row.Status)
}
}

View File

@@ -1,176 +0,0 @@
package batchtfile
import (
"context"
"encoding/json"
"fmt"
"path/filepath"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/core"
tftask "github.com/krau/SaveAny-Bot/core/tasks/tfile"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
"github.com/krau/SaveAny-Bot/pkg/enums/tasktype"
tfilepkg "github.com/krau/SaveAny-Bot/pkg/tfile"
"github.com/krau/SaveAny-Bot/storage"
)
type elementPayload struct {
ID string `json:"id"`
Storage string `json:"storage"`
Path string `json:"path"`
File tfilepkg.FilePayload `json:"file"`
SourceGroupKey string `json:"source_group_key"`
SourceCaption string `json:"source_caption"`
PreserveCaption bool `json:"preserve_caption"`
}
type taskPayload struct {
Kind string `json:"kind"` // "batch"
ID string `json:"id"`
Elements []elementPayload `json:"elements"`
ChatID int64 `json:"chat_id"`
MessageID int `json:"message_id"`
IgnoreErrors bool `json:"ignore_errors"`
Overwrite bool `json:"overwrite"`
// Done lists element IDs whose upload completed; they are skipped on recovery.
Done []string `json:"done"`
}
// tgfilesCodec is the single codec registered for TaskTypeTgfiles: it
// dispatches between single-file and batch tasks by concrete type on marshal
// and by payload shape on unmarshal. Registering one codec per task class
// under the shared TaskTypeTgfiles key would let the last init() win and
// silently disable persistence for the other class.
type tgfilesCodec struct{}
func init() {
core.RegisterTaskCodec(tasktype.TaskTypeTgfiles, tgfilesCodec{})
}
func (tgfilesCodec) Marshal(task core.Executable) ([]byte, error) {
switch t := task.(type) {
case *tftask.Task:
return tftask.TaskCodec.Marshal(t)
case *Task:
return batchCodec{}.Marshal(t)
default:
return nil, fmt.Errorf("unexpected task type %T", task)
}
}
// detectTaskKind returns "batch" or "file" for a persisted tgfiles payload.
// New payloads carry an explicit kind; legacy payloads are detected by shape.
func detectTaskKind(data []byte) (string, error) {
var shape struct {
Kind string `json:"kind"`
Elements []json.RawMessage `json:"elements"`
File json.RawMessage `json:"file"`
}
if err := json.Unmarshal(data, &shape); err != nil {
return "", fmt.Errorf("invalid task payload: %w", err)
}
switch {
case shape.Kind == "batch", shape.Kind == "" && shape.Elements != nil:
return "batch", nil
case shape.Kind == "file", shape.Kind == "" && shape.File != nil:
return "file", nil
default:
return "", fmt.Errorf("unrecognized task payload")
}
}
func (tgfilesCodec) Unmarshal(data []byte) (core.Executable, error) {
kind, err := detectTaskKind(data)
if err != nil {
return nil, err
}
if kind == "batch" {
return batchCodec{}.Unmarshal(data)
}
return tftask.TaskCodec.Unmarshal(data)
}
type batchCodec struct{}
func (batchCodec) Marshal(task core.Executable) ([]byte, error) {
t, ok := task.(*Task)
if !ok {
return nil, fmt.Errorf("unexpected task type %T", task)
}
p := taskPayload{
Kind: "batch",
ID: t.ID,
IgnoreErrors: t.IgnoreErrors,
Done: t.completedElementIDs(),
}
if overwrite, ok := t.ctx.Value(ctxkey.OverwriteExisting).(bool); ok {
p.Overwrite = overwrite
}
for _, elem := range t.elems {
filePayload, ok := tfilepkg.FilePayloadOf(elem.File)
if !ok {
return nil, fmt.Errorf("file %T is not serializable", elem.File)
}
p.Elements = append(p.Elements, elementPayload{
ID: elem.ID,
Storage: elem.Storage.Name(),
Path: elem.Path,
File: filePayload,
SourceGroupKey: elem.sourceGroupKey,
SourceCaption: elem.sourceCaption,
PreserveCaption: elem.preserveCaption,
})
}
if progress, ok := t.Progress.(*Progress); ok {
p.ChatID = progress.ChatID
p.MessageID = progress.MessageID
}
return json.Marshal(p)
}
func (batchCodec) Unmarshal(data []byte) (core.Executable, error) {
var p taskPayload
if err := json.Unmarshal(data, &p); err != nil {
return nil, fmt.Errorf("invalid task payload: %w", err)
}
dler := core.DownloaderClient()
if dler == nil {
return nil, fmt.Errorf("no downloader client available")
}
done := make(map[string]struct{}, len(p.Done))
for _, id := range p.Done {
done[id] = struct{}{}
}
elems := make([]TaskElement, 0, len(p.Elements))
for _, ep := range p.Elements {
if _, ok := done[ep.ID]; ok {
continue // upload already completed; do not re-run
}
stor, err := storage.GetStorageByName(context.Background(), ep.Storage)
if err != nil {
return nil, fmt.Errorf("storage %q: %w", ep.Storage, err)
}
localPath, err := filepath.Abs(filepath.Join(config.C().Temp.BasePath, fmt.Sprintf("%s_%s", ep.ID, ep.File.Name)))
if err != nil {
return nil, fmt.Errorf("failed to build cache path: %w", err)
}
elems = append(elems, TaskElement{
ID: ep.ID,
Storage: stor,
Path: ep.Path,
File: tfilepkg.FileFromPayload(ep.File, dler),
localPath: localPath,
sourceGroupKey: ep.SourceGroupKey,
sourceCaption: ep.SourceCaption,
preserveCaption: ep.PreserveCaption,
})
}
var progress ProgressTracker
if p.ChatID != 0 {
progress = NewProgressTracker(p.MessageID, p.ChatID)
}
task := NewBatchTGFileTask(p.ID, context.Background(), elems, progress, p.IgnoreErrors)
task.overwrite = p.Overwrite
return task, nil
}

View File

@@ -1,39 +0,0 @@
package batchtfile
import (
"testing"
)
func TestDetectTaskKind(t *testing.T) {
tests := []struct {
name string
payload string
want string
wantErr bool
}{
{"batch with kind", `{"kind":"batch","id":"1","elements":[]}`, "batch", false},
{"file with kind", `{"kind":"file","id":"1","file":{}}`, "file", false},
{"legacy batch by shape", `{"id":"1","elements":[]}`, "batch", false},
{"legacy file by shape", `{"id":"1","file":{}}`, "file", false},
{"legacy batch with element", `{"id":"1","elements":[{"id":"e"}]}`, "batch", false},
{"no discriminator", `{"id":"1"}`, "", true},
{"invalid json", `not json`, "", true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := detectTaskKind([]byte(tt.payload))
if tt.wantErr {
if err == nil {
t.Fatalf("expected error, got kind %q", got)
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != tt.want {
t.Fatalf("kind = %q, want %q", got, tt.want)
}
})
}
}

View File

@@ -1,92 +0,0 @@
package batchtfile
import (
"context"
"fmt"
"os"
"time"
"github.com/charmbracelet/log"
"github.com/krau/SaveAny-Bot/common/tdler"
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
"github.com/krau/SaveAny-Bot/common/utils/fsutil"
"github.com/krau/SaveAny-Bot/common/utils/ioutil"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/pkg/taskevent"
)
// downloadToCache fetches elem.File into the element cache path. It resumes
// from a partial .part download tracked by a resume bitmap, and reuses a
// complete cache file (e.g. when the previous run was interrupted during
// upload).
func (t *Task) downloadToCache(ctx context.Context, elem *TaskElement) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", elem.File.Name()))
if elem.File.Size() > 0 {
if stat, err := os.Stat(elem.localPath); err == nil && stat.Size() == elem.File.Size() {
logger.Info("Cache file already complete, skipping download")
return nil
}
}
onProgress := t.downloadCallback(ctx, elem)
if elem.File.Size() <= 0 {
// Unknown size (e.g. photos) cannot be resumed; use the plain downloader.
localFile, err := fsutil.CreateFile(elem.localPath)
if err != nil {
return fmt.Errorf("failed to create local file: %w", err)
}
defer localFile.Close()
wrAt := ioutil.NewProgressWriterAt(localFile, onProgress)
if _, err := tdler.NewDownloader(elem.File).Parallel(ctx, wrAt); err != nil {
return err
}
return nil
}
partPath := elem.localPath + ".part"
// 不截断已存在的 .part: 位图标记的已完成块依赖既有字节。
localFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644)
if err != nil {
return fmt.Errorf("failed to create local file: %w", err)
}
wrAt := ioutil.NewProgressWriterAt(localFile, onProgress)
err = tdler.DownloadResumable(
ctx, elem.File, wrAt,
dlutil.BestThreads(elem.File.Size(), config.C().Threads),
tdler.ResumeStatePath(partPath),
)
closeErr := localFile.Close()
if err != nil {
return err
}
if closeErr != nil {
return fmt.Errorf("failed to close cache file: %w", closeErr)
}
stat, err := os.Stat(partPath)
if err != nil {
return fmt.Errorf("failed to stat downloaded file: %w", err)
}
if stat.Size() != elem.File.Size() {
return fmt.Errorf("downloaded size %d does not match expected %d", stat.Size(), elem.File.Size())
}
if err := os.Rename(partPath, elem.localPath); err != nil {
return fmt.Errorf("failed to finalize download: %w", err)
}
// 清理位图是尽力而为: 下载已完成, 清理失败不应使任务失败。
if err := tdler.RemoveResumeState(tdler.ResumeStatePath(partPath)); err != nil {
logger.Warnf("Failed to remove resume state: %v", err)
}
return nil
}
func (t *Task) downloadCallback(ctx context.Context, elem *TaskElement) func(int) {
return func(n int) {
t.recordItemDownload(elem.ID, int64(n), time.Now())
downloaded := t.downloaded.Add(int64(n))
t.notifyProgress(ctx)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
TotalBytes: t.totalSize,
DownloadedBytes: downloaded,
})
}
}

View File

@@ -2,13 +2,10 @@ package batchtfile
import (
"context"
"errors"
"fmt"
"io"
"os"
"path"
"sync"
"time"
"github.com/charmbracelet/log"
"github.com/duke-git/lancet/v2/retry"
@@ -17,309 +14,45 @@ import (
"github.com/krau/SaveAny-Bot/common/utils/ioutil"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
"github.com/krau/SaveAny-Bot/pkg/storagetypes"
"github.com/krau/SaveAny-Bot/pkg/taskevent"
"github.com/krau/SaveAny-Bot/storage"
"golang.org/x/sync/errgroup"
)
type executionGroup struct {
elems []*TaskElement
batchSaver storage.StorageBatchSaver
}
func (g executionGroup) usesBatchSaver() bool {
return g.batchSaver != nil
}
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")
if t.overwrite {
ctx = storage.WithOverwrite(ctx)
}
if t.Progress != nil {
t.Progress.OnStart(ctx, t)
}
groups := t.executionGroups()
var err error
for i := 0; i < len(groups); {
if groups[i].usesBatchSaver() {
err = t.processBatch(ctx, groups[i])
i++
} else {
end := i + 1
for end < len(groups) && !groups[end].usesBatchSaver() {
end++
t.Progress.OnStart(ctx, t)
workers := config.C().Workers
eg, gctx := errgroup.WithContext(ctx)
eg.SetLimit(workers)
for _, elem := range t.elems {
eg.Go(func() error {
t.processingMu.RLock()
if t.processing[elem.ID] != nil {
return fmt.Errorf("element with ID %s is already being processed", elem.ID)
}
elems := make([]*TaskElement, 0, end-i)
for _, group := range groups[i:end] {
elems = append(elems, group.elems...)
}
err = t.processElements(ctx, elems)
i = end
}
if err != nil {
if !t.IgnoreErrors || errors.Is(err, context.Canceled) {
break
}
logger.Warnf("Group processing failed (ignored): %v", err)
err = nil
}
t.processingMu.RUnlock()
t.processingMu.Lock()
t.processing[elem.ID] = &elem
t.processingMu.Unlock()
defer func() {
t.processingMu.Lock()
delete(t.processing, elem.ID)
t.processingMu.Unlock()
}()
return t.processElement(gctx, elem)
})
}
err := eg.Wait()
if err != nil {
logger.Errorf("Error during batch file processing: %v", err)
} else {
logger.Info("Batch file task completed successfully")
}
t.finishItems(err)
if t.Progress != nil {
t.Progress.OnDone(ctx, t, err)
}
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); {
elem := &t.elems[i]
batchSaver, batchCapable := elem.Storage.(storage.StorageBatchSaver)
if !batchCapable || elem.sourceGroupKey == "" {
groups = append(groups, executionGroup{elems: []*TaskElement{elem}})
i++
continue
}
end := i + 1
for end < len(t.elems) {
next := &t.elems[end]
if next.Storage != elem.Storage || next.sourceGroupKey != elem.sourceGroupKey {
break
}
end++
}
elems := make([]*TaskElement, 0, end-i)
for j := i; j < end; j++ {
elems = append(elems, &t.elems[j])
}
groups = append(groups, executionGroup{elems: elems, batchSaver: batchSaver})
i = end
}
return groups
}
func (t *Task) processElements(ctx context.Context, elems []*TaskElement) error {
eg, gctx := errgroup.WithContext(ctx)
eg.SetLimit(config.C().Workers)
for _, elem := range elems {
eg.Go(func() error {
if err := t.markProcessing(ctx, elem); err != nil {
return err
}
defer t.unmarkProcessing(elem.ID)
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()
}
func (t *Task) processBatch(ctx context.Context, group executionGroup) error {
// Cache files are kept on failure so a later restart can resume upload.
uploaded := false
defer func() {
if !uploaded {
return
}
for _, elem := range group.elems {
if err := os.Remove(elem.localPath); err != nil && !os.IsNotExist(err) {
log.FromContext(ctx).Warnf("Failed to cleanup batch cache file %s: %v", elem.localPath, err)
}
}
}()
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 i, elem := range group.elems {
eg.Go(func() error {
if err := t.markProcessing(ctx, elem); err != nil {
return err
}
defer t.unmarkProcessing(elem.ID)
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
}
// 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 {
log.FromContext(ctx).Warnf("Failed to close batch cache file %s: %v", file.Name(), err)
}
}
}()
for _, elem := range successElems {
file, err := os.Open(elem.localPath)
if err != nil {
t.markItemFailed(elem.ID, FailureStageCache, err)
return fmt.Errorf("failed to open cache file: %w", err)
}
stat, err := file.Stat()
if err != nil {
file.Close()
t.markItemFailed(elem.ID, FailureStageCache, err)
return fmt.Errorf("failed to get cache file stat: %w", err)
}
openFiles = append(openFiles, file)
items = append(items, storagetypes.BatchItem{
Reader: file,
StoragePath: elem.Path,
Size: stat.Size(),
SourceGroupKey: elem.sourceGroupKey,
Caption: elem.sourceCaption,
PreserveCaption: elem.preserveCaption,
})
}
for index, item := range items {
t.recordDownloadComplete(successElems[index].ID, item.Size)
}
err := t.saveBatchItems(ctx, successElems, items)
if err == nil {
uploaded = true
for _, elem := range successElems {
t.persistElementDone(ctx, elem.ID)
}
}
return err
}
func (t *Task) saveBatchItems(ctx context.Context, successElems []*TaskElement, items []storagetypes.BatchItem) error {
t.startUpload(ctx)
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(successElems) {
return
}
t.uploadCallback(ctx, successElems[index].ID)(uploaded, total)
})
if 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 successElems {
t.uploadCallback(ctx, elem.ID)(items[index].Size, items[index].Size)
t.markItemCompleted(elem.ID)
}
t.notifyStateChange(ctx)
return nil
}
for i := range items {
items[i].Reader = ioutil.NewProgressReader(
items[i].Reader,
items[i].Size,
t.uploadCallback(ctx, successElems[i].ID),
)
}
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 successElems {
t.uploadCallback(ctx, elem.ID)(items[index].Size, items[index].Size)
t.markItemCompleted(elem.ID)
}
t.notifyStateChange(ctx)
return nil
}
func (t *Task) markProcessing(ctx context.Context, elem *TaskElement) error {
t.processingMu.Lock()
if t.processing[elem.ID] != nil {
t.processingMu.Unlock()
return fmt.Errorf("element with ID %s is already being processed", elem.ID)
}
t.processing[elem.ID] = elem
t.processingMu.Unlock()
t.markItemActive(elem.ID, elem.stream, time.Now())
t.notifyProgress(ctx)
return nil
}
func (t *Task) unmarkProcessing(id string) {
t.processingMu.Lock()
delete(t.processing, id)
t.processingMu.Unlock()
}
func (t *Task) downloadElement(ctx context.Context, elem *TaskElement) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", elem.File.Name()))
logger.Info("Starting file download")
if err := t.downloadToCache(ctx, elem); err != nil {
t.markItemFailed(elem.ID, FailureStageDownload, err)
t.notifyStateChange(ctx)
return fmt.Errorf("failed to download file: %w", err)
}
logger.Info("File downloaded successfully")
if path.Ext(elem.FileName()) == "" {
if ext := fsutil.DetectFileExt(elem.localPath); ext != "" {
elem.Path += ext
}
}
t.markItemDownloaded(elem.ID)
t.notifyProgress(ctx)
return nil
}
func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", elem.File.Name()))
if elem.stream {
@@ -327,17 +60,11 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
defer pr.Close()
errg, uploadCtx := errgroup.WithContext(ctx)
errg.Go(func() error {
err := elem.Storage.Save(uploadCtx, pr, elem.Path)
if err != nil {
t.markItemFailed(elem.ID, FailureStageUpload, err)
t.notifyStateChange(ctx)
}
return err
return elem.Storage.Save(uploadCtx, pr, elem.Path)
})
wr := ioutil.NewProgressWriter(pw, func(n int) {
t.recordItemDownload(elem.ID, int64(n), time.Now())
downloaded := t.downloaded.Add(int64(n))
t.notifyProgress(ctx)
t.Progress.OnProgress(ctx, t)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
@@ -351,8 +78,6 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
_, err := tdler.NewDownloader(elem.File).Stream(uploadCtx, wr)
if err != nil {
logger.Errorf("Failed to download file: %v", err)
t.markItemFailed(elem.ID, FailureStageDownload, err)
t.notifyStateChange(ctx)
pw.CloseWithError(err)
}
return err
@@ -360,30 +85,31 @@ 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)
}
// 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")
return nil
}
logger.Info("Starting file download")
// 不预创建缓存文件: 预创建会截断上次运行保留的完整缓存, 使复用失效。
success := false
localFile, err := fsutil.CreateFile(elem.localPath)
if err != nil {
return fmt.Errorf("failed to create local file: %w", err)
}
defer func() {
if success {
if err := os.Remove(elem.localPath); err != nil {
logger.Errorf("Failed to remove cache file: %v", err)
}
if err := localFile.CloseAndRemove(); err != nil {
logger.Errorf("Failed to close local file: %v", err)
}
}()
if err := t.downloadToCache(ctx, &elem); err != nil {
t.markItemFailed(elem.ID, FailureStageDownload, err)
t.notifyStateChange(ctx)
wrAt := ioutil.NewProgressWriterAt(localFile, func(n int) {
downloaded := t.downloaded.Add(int64(n))
t.Progress.OnProgress(ctx, t)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
TotalBytes: t.totalSize,
DownloadedBytes: downloaded,
})
})
_, err = tdler.NewDownloader(elem.File).Parallel(ctx, wrAt)
if err != nil {
return fmt.Errorf("failed to download file: %w", err)
}
logger.Info("File downloaded successfully")
@@ -393,54 +119,24 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
elem.Path = elem.Path + ext
}
}
fileStat, err := os.Stat(elem.localPath)
var fileStat os.FileInfo
fileStat, err = os.Stat(elem.localPath)
if err != nil {
t.markItemFailed(elem.ID, FailureStageCache, err)
t.notifyStateChange(ctx)
return fmt.Errorf("failed to get file stat: %w", err)
}
t.recordDownloadComplete(elem.ID, fileStat.Size())
vctx := context.WithValue(ctx, ctxkey.ContentLength, fileStat.Size())
t.startUpload(vctx)
onProgress := t.uploadCallback(vctx, elem.ID)
attempt := 0
retryLimit := int(config.C().Retry)
lastFailureStage := FailureStageUpload
err = retry.Retry(func() error {
attempt++
var file *os.File
file, err = os.Open(elem.localPath)
if err != nil {
lastFailureStage = FailureStageCache
t.markItemRetry(elem.ID, lastFailureStage, attempt, retryLimit, err)
t.notifyStateChange(vctx)
return fmt.Errorf("failed to open cache file: %w", err)
}
defer file.Close()
onProgress(0, fileStat.Size())
if progressSaver, ok := elem.Storage.(storage.StorageProgressSaver); ok {
err = progressSaver.SaveWithProgress(vctx, file, elem.Path, onProgress)
} else {
err = elem.Storage.Save(vctx, ioutil.NewProgressReader(file, fileStat.Size(), onProgress), elem.Path)
}
if err != nil {
if err = elem.Storage.Save(vctx, file, elem.Path); err != nil {
logger.Errorf("Failed to save file: %s, retrying...", err)
lastFailureStage = t.itemFailureStage(elem.ID)
t.markItemRetry(elem.ID, lastFailureStage, attempt, retryLimit, err)
t.notifyStateChange(vctx)
return err
}
return nil
}, retry.Context(vctx), retry.RetryTimes(uint(config.C().Retry)))
if err == nil {
onProgress(fileStat.Size(), fileStat.Size())
t.markItemCompleted(elem.ID)
t.notifyStateChange(vctx)
t.persistElementDone(ctx, elem.ID)
success = true
} else {
t.markItemFailed(elem.ID, lastFailureStage, err)
t.notifyStateChange(vctx)
}
return err
}

View File

@@ -1,57 +0,0 @@
package batchtfile
import (
"testing"
"github.com/gotd/td/tg"
"github.com/krau/SaveAny-Bot/pkg/tfile"
tgstorage "github.com/krau/SaveAny-Bot/storage/telegram"
)
func TestExecutionGroupsPreserveSourceAlbums(t *testing.T) {
stor := new(tgstorage.Telegram)
otherStor := new(tgstorage.Telegram)
task := Task{elems: []TaskElement{
{Storage: stor, sourceGroupKey: "album-1"},
{Storage: stor, sourceGroupKey: "album-1"},
{Storage: stor},
{Storage: stor, sourceGroupKey: "album-2"},
{Storage: stor, sourceGroupKey: "album-2"},
{Storage: otherStor, sourceGroupKey: "album-2"},
}}
groups := task.executionGroups()
wantSizes := []int{2, 1, 2, 1}
wantBatch := []bool{true, false, true, true}
if len(groups) != len(wantSizes) {
t.Fatalf("got %d groups, want %d", len(groups), len(wantSizes))
}
for i := range groups {
if got := len(groups[i].elems); got != wantSizes[i] {
t.Errorf("group %d has %d elements, want %d", i, got, wantSizes[i])
}
if got := groups[i].usesBatchSaver(); got != wantBatch[i] {
t.Errorf("group %d batch=%v, want %v", i, got, wantBatch[i])
}
}
}
func TestSourceMetadataPreservesAlbumIdentityAndCaption(t *testing.T) {
msg := &tg.Message{
PeerID: &tg.PeerChannel{ChannelID: 77},
Message: "original caption",
}
msg.SetGroupedID(42)
file := tfile.NewTGFile(nil, nil, 0, "photo.jpg", tfile.WithMessage(msg))
groupKey, caption, preserveCaption := sourceMetadata(file)
if groupKey != "*tg.PeerChannel:77:42" {
t.Fatalf("group key = %q, want %q", groupKey, "*tg.PeerChannel:77:42")
}
if caption != "original caption" {
t.Fatalf("caption = %q, want original caption", caption)
}
if !preserveCaption {
t.Fatal("preserveCaption = false, want true")
}
}

View File

@@ -1,360 +0,0 @@
package batchtfile
import (
"context"
"errors"
"strings"
"time"
)
const (
transferSpeedWindow = 5 * time.Second
transferSamplePeriod = 250 * time.Millisecond
)
// ItemPhase describes the current lifecycle stage of one batch item.
type ItemPhase uint8
const (
ItemPhaseWaiting ItemPhase = iota
ItemPhaseDownloading
ItemPhaseTransferring
ItemPhaseDownloaded
ItemPhaseUploading
ItemPhaseRetrying
ItemPhaseConfirming
ItemPhaseCompleted
ItemPhaseFailed
ItemPhaseStopped
)
// FailureStage identifies the operation that failed for one batch item.
type FailureStage uint8
const (
FailureStageNone FailureStage = iota
FailureStageDownload
FailureStageCache
FailureStageUpload
FailureStageConfirm
FailureStageBatchUpload
FailureStageInternal
)
// TaskItemProgress is an immutable progress snapshot for one batch item.
type TaskItemProgress struct {
Index int
ID string
Name string
Size int64
Downloaded int64
Uploaded int64
DownloadSpeed float64
UploadSpeed float64
Phase ItemPhase
FailureStage FailureStage
RetryAttempt int
RetryLimit int
Error string
}
type transferSample struct {
at time.Time
bytes int64
}
type transferMeter struct {
samples []transferSample
latest transferSample
hasData bool
}
func (m *transferMeter) record(now time.Time, transferred int64) {
if m.hasData && transferred < m.latest.bytes {
m.reset()
}
if m.hasData && now.Before(m.latest.at) {
now = m.latest.at
}
sample := transferSample{at: now, bytes: transferred}
m.latest = sample
m.hasData = true
if len(m.samples) == 0 {
m.samples = append(m.samples, sample)
return
}
if now.Sub(m.samples[len(m.samples)-1].at) >= transferSamplePeriod {
m.samples = append(m.samples, sample)
}
cutoff := now.Add(-transferSpeedWindow)
for len(m.samples) > 1 && m.samples[0].at.Before(cutoff) {
m.samples = m.samples[1:]
}
}
func (m *transferMeter) speed() float64 {
if !m.hasData || len(m.samples) == 0 {
return 0
}
first := m.samples[0]
last := m.latest
elapsed := last.at.Sub(first.at).Seconds()
if elapsed <= 0 || last.bytes <= first.bytes {
return 0
}
return float64(last.bytes-first.bytes) / elapsed
}
func (m *transferMeter) reset() {
m.samples = m.samples[:0]
m.latest = transferSample{}
m.hasData = false
}
type itemProgressState struct {
index int
id string
name string
expectedSize int64
actualSize int64
downloaded int64
uploaded int64
phase ItemPhase
failureStage FailureStage
retryAttempt int
retryLimit int
err string
downloadMeter transferMeter
uploadMeter transferMeter
}
func newItemProgressStates(elems []TaskElement) ([]itemProgressState, map[string]int) {
states := make([]itemProgressState, 0, len(elems))
index := make(map[string]int, len(elems))
for i, elem := range elems {
name := ""
size := int64(0)
if elem.File != nil {
name = elem.File.Name()
size = elem.File.Size()
}
states = append(states, itemProgressState{
index: i + 1,
id: elem.ID,
name: name,
expectedSize: size,
phase: ItemPhaseWaiting,
})
index[elem.ID] = i
}
return states, index
}
func (t *Task) updateItem(id string, update func(*itemProgressState)) bool {
t.itemMu.Lock()
defer t.itemMu.Unlock()
index, ok := t.itemIndex[id]
if !ok || index < 0 || index >= len(t.itemStates) {
return false
}
update(&t.itemStates[index])
return true
}
func (t *Task) markItemActive(id string, stream bool, now time.Time) {
t.updateItem(id, func(item *itemProgressState) {
if stream {
item.phase = ItemPhaseTransferring
} else {
item.phase = ItemPhaseDownloading
}
item.failureStage = FailureStageNone
item.err = ""
item.downloadMeter.record(now, item.downloaded)
if stream {
item.uploadMeter.record(now, item.uploaded)
}
})
}
func (t *Task) recordItemDownload(id string, n int64, now time.Time) {
if n <= 0 {
return
}
t.updateItem(id, func(item *itemProgressState) {
item.downloaded += n
item.downloadMeter.record(now, item.downloaded)
if item.phase == ItemPhaseTransferring {
item.uploaded += n
item.uploadMeter.record(now, item.uploaded)
}
})
}
func (t *Task) markItemDownloaded(id string) {
t.updateItem(id, func(item *itemProgressState) {
item.phase = ItemPhaseDownloaded
})
}
func (t *Task) recordItemDownloaded(id string, actualSize int64) {
t.updateItem(id, func(item *itemProgressState) {
if actualSize > 0 {
item.actualSize = actualSize
}
if item.actualSize == 0 {
item.actualSize = item.downloaded
}
if item.phase != ItemPhaseTransferring {
item.phase = ItemPhaseDownloaded
item.uploadMeter.reset()
}
})
}
func (t *Task) recordItemUpload(id string, uploaded, total int64, now time.Time) bool {
becameConfirming := false
t.updateItem(id, func(item *itemProgressState) {
if uploaded < item.uploaded {
if item.phase != ItemPhaseRetrying {
return
}
item.uploadMeter.reset()
}
if total > 0 {
item.actualSize = total
}
item.uploaded = uploaded
item.uploadMeter.record(now, uploaded)
if total > 0 && uploaded >= total {
becameConfirming = item.phase != ItemPhaseConfirming
item.phase = ItemPhaseConfirming
return
}
item.phase = ItemPhaseUploading
item.failureStage = FailureStageNone
item.err = ""
})
return becameConfirming
}
func (t *Task) markItemRetry(id string, stage FailureStage, attempt, limit int, err error) {
t.updateItem(id, func(item *itemProgressState) {
item.phase = ItemPhaseRetrying
item.failureStage = stage
item.retryAttempt = attempt
item.retryLimit = limit
item.err = compactError(err)
})
}
func (t *Task) markItemFailed(id string, stage FailureStage, err error) {
t.updateItem(id, func(item *itemProgressState) {
if item.phase == ItemPhaseFailed || item.phase == ItemPhaseCompleted {
return
}
if errors.Is(err, context.Canceled) {
item.phase = ItemPhaseStopped
return
}
item.phase = ItemPhaseFailed
item.failureStage = stage
item.err = compactError(err)
})
}
func (t *Task) markItemCompleted(id string) {
t.updateItem(id, func(item *itemProgressState) {
item.phase = ItemPhaseCompleted
item.failureStage = FailureStageNone
item.err = ""
item.retryAttempt = 0
item.retryLimit = 0
if item.actualSize == 0 {
item.actualSize = max(item.downloaded, item.uploaded)
}
})
}
func (t *Task) finishItems(err error) {
if err == nil {
return
}
t.itemMu.Lock()
defer t.itemMu.Unlock()
for i := range t.itemStates {
item := &t.itemStates[i]
if item.phase == ItemPhaseCompleted || item.phase == ItemPhaseFailed {
continue
}
item.phase = ItemPhaseStopped
}
}
func (t *Task) itemFailureStage(id string) FailureStage {
t.itemMu.RLock()
defer t.itemMu.RUnlock()
index, ok := t.itemIndex[id]
if !ok || index < 0 || index >= len(t.itemStates) {
return FailureStageUpload
}
if t.itemStates[index].phase == ItemPhaseConfirming {
return FailureStageConfirm
}
return FailureStageUpload
}
func (t *Task) Items() []TaskItemProgress {
t.itemMu.RLock()
defer t.itemMu.RUnlock()
items := make([]TaskItemProgress, 0, len(t.itemStates))
for i := range t.itemStates {
item := &t.itemStates[i]
size := item.actualSize
if size == 0 {
size = item.expectedSize
}
items = append(items, TaskItemProgress{
Index: item.index,
ID: item.id,
Name: item.name,
Size: size,
Downloaded: item.downloaded,
Uploaded: item.uploaded,
DownloadSpeed: item.downloadMeter.speed(),
UploadSpeed: item.uploadMeter.speed(),
Phase: item.phase,
FailureStage: item.failureStage,
RetryAttempt: item.retryAttempt,
RetryLimit: item.retryLimit,
Error: item.err,
})
}
return items
}
func (t *Task) ActualTotalSize() int64 {
items := t.Items()
var total int64
for _, item := range items {
total += item.Size
}
return total
}
type stateProgressTracker interface {
OnStateChange(ctx context.Context, info TaskInfo)
}
func (t *Task) notifyStateChange(ctx context.Context) {
if tracker, ok := t.Progress.(stateProgressTracker); ok {
tracker.OnStateChange(ctx, t)
}
}
func compactError(err error) string {
if err == nil {
return ""
}
return strings.Join(strings.Fields(err.Error()), " ")
}

View File

@@ -1,35 +0,0 @@
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
}

View File

@@ -4,19 +4,20 @@ import (
"context"
"errors"
"fmt"
"path"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"unicode/utf8"
"github.com/charmbracelet/log"
"github.com/duke-git/lancet/v2/slice"
"github.com/gotd/td/telegram/message/entity"
"github.com/gotd/td/telegram/message/styling"
"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/dlutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
"github.com/krau/SaveAny-Bot/config"
)
type ProgressTracker interface {
@@ -26,475 +27,159 @@ type ProgressTracker interface {
}
type Progress struct {
MessageID int
ChatID int64
updateMu sync.Mutex
lastUpdateAt time.Time
lastText string
done bool
skippedFiles []string
MessageID int
ChatID int64
start time.Time
lastUpdatePercent atomic.Int32
skippedFiles []string
}
type renderedBatchMessage struct {
Text string
Entities []tg.MessageEntityClass
Err error
}
const (
progressRenderInterval = time.Second
maxVisibleActiveItems = 5
progressBarWidth = 10
maxDisplayNameRunes = 36
maxDisplayErrorRunes = 240
)
func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
p.render(ctx, info, true)
p.start = time.Now()
p.lastUpdatePercent.Store(0)
log.FromContext(ctx).Debugf("Batch task progress tracking started for message %d in chat %d", p.MessageID, p.ChatID)
entityBuilder := entity.Builder{}
var entities []tg.MessageEntityClass
if err := styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressBatchStartPrefix, nil)),
styling.Code(fmt.Sprintf("%.2f MB (%d个文件)", float64(info.TotalSize())/(1024*1024), info.Count())),
); err != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", err)
return
}
text, entities := entityBuilder.Complete()
req := &tg.MessagesEditMessageRequest{
ID: p.MessageID,
}
req.SetMessage(text)
req.SetEntities(entities)
req.SetReplyMarkup(&tg.ReplyInlineMarkup{
Rows: []tg.KeyboardButtonRow{
{
Buttons: []tg.KeyboardButtonClass{
tgutil.BuildCancelButton(info.TaskID()),
},
},
}},
)
ext := tgutil.ExtFromContext(ctx)
if ext != nil {
ext.EditMessage(p.ChatID, req)
return
}
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
p.render(ctx, info, false)
}
func (p *Progress) OnStateChange(ctx context.Context, info TaskInfo) {
p.render(ctx, info, true)
}
func (p *Progress) OnUploadStart(ctx context.Context, info TaskInfo, _ int64) {
p.render(ctx, info, true)
}
func (p *Progress) OnUploadProgress(ctx context.Context, info TaskInfo, _, _ int64) {
p.render(ctx, info, false)
}
func (p *Progress) render(ctx context.Context, info TaskInfo, priority bool) {
p.updateMu.Lock()
defer p.updateMu.Unlock()
if p.done {
if !shouldUpdateProgress(info.TotalSize(), info.Downloaded(), int(p.lastUpdatePercent.Load())) {
return
}
now := time.Now()
if !priority && !p.lastUpdateAt.IsZero() && now.Sub(p.lastUpdateAt) < progressRenderInterval {
percent := int((info.Downloaded() * 100) / info.TotalSize())
if p.lastUpdatePercent.Load() == int32(percent) {
return
}
message := buildBatchProgressMessage(info, p.skippedFiles, visibleActiveItems())
if message.Err != nil {
log.FromContext(ctx).Errorf("Failed to render batch progress message: %v", message.Err)
p.lastUpdatePercent.Store(int32(percent))
log.FromContext(ctx).Debugf("Progress update: %s, %d/%d", info.TaskID(), info.Downloaded(), info.TotalSize())
entityBuilder := entity.Builder{}
var entities []tg.MessageEntityClass
if err := styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressBatchProcessingPrefix, nil)),
styling.Code(fmt.Sprintf("%.2f MB (%d个文件)", float64(info.TotalSize())/(1024*1024), info.Count())),
styling.Plain(i18n.T(i18nk.BotMsgProgressProcessingListPrefix, nil)),
func() styling.StyledTextOption {
var lines []string
for _, elem := range info.Processing() {
lines = append(lines, fmt.Sprintf(" - %s (%.2f MB)", elem.FileName(), float64(elem.FileSize())/(1024*1024)))
}
if len(lines) == 0 {
lines = append(lines, i18n.T(i18nk.BotMsgProgressProcessingNone, nil))
}
return styling.Plain(slice.Join(lines, "\n"))
}(),
styling.Plain(i18n.T(i18nk.BotMsgProgressAvgSpeedPrefix, nil)),
styling.Bold(fmt.Sprintf("%.2f MB/s", dlutil.GetSpeed(info.Downloaded(), p.start)/(1024*1024))),
styling.Plain(i18n.T(i18nk.BotMsgProgressCurrentProgressPrefix, nil)),
styling.Bold(fmt.Sprintf("%.2f%%", float64(info.Downloaded())/float64(info.TotalSize())*100)),
); err != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", err)
return
}
if message.Text == p.lastText {
text, entities := entityBuilder.Complete()
req := &tg.MessagesEditMessageRequest{
ID: p.MessageID,
}
req.SetMessage(text)
req.SetEntities(entities)
req.SetReplyMarkup(&tg.ReplyInlineMarkup{
Rows: []tg.KeyboardButtonRow{
{
Buttons: []tg.KeyboardButtonClass{
tgutil.BuildCancelButton(info.TaskID()),
},
},
}},
)
ext := tgutil.ExtFromContext(ctx)
if ext != nil {
ext.EditMessage(p.ChatID, req)
return
}
p.lastText = message.Text
p.lastUpdateAt = now
p.editMessage(ctx, info.TaskID(), message, true)
}
func (p *Progress) OnDone(ctx context.Context, info TaskInfo, err error) {
p.updateMu.Lock()
defer p.updateMu.Unlock()
if p.done {
if err != nil {
log.FromContext(ctx).Errorf("Batch task %s failed: %s", info.TaskID(), err)
} else {
log.FromContext(ctx).Debugf("Batch task %s completed successfully", info.TaskID())
}
entityBuilder := entity.Builder{}
var stylingErr error
if err != nil {
if errors.Is(err, context.Canceled) {
stylingErr = styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressTaskCanceled, nil)),
)
} else {
stylingErr = styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressTaskFailedWithError, map[string]any{
"Error": "",
})),
styling.Code(err.Error()),
)
}
} else {
stylingErr = styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressBatchDonePrefix, nil)),
styling.Code(strconv.Itoa(info.Count())),
styling.Plain(i18n.T(i18nk.BotMsgProgressTotalSizePrefix, nil)),
styling.Code(fmt.Sprintf("%.2f MB", float64(info.TotalSize())/(1024*1024))),
func() styling.StyledTextOption {
if len(p.skippedFiles) == 0 {
return styling.Plain("")
}
return styling.Plain("\n\n" + i18n.T(i18nk.BotMsgCommonInfoConflictFilesSkipped, map[string]any{
"Skipped": strings.Join(p.skippedFiles, "\n"),
}))
}(),
)
}
if stylingErr != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", stylingErr)
return
}
p.done = true
message := buildBatchDoneMessage(info, p.skippedFiles, err)
if message.Err != nil {
log.FromContext(ctx).Errorf("Failed to render final batch progress message: %v", message.Err)
return
}
p.lastText = message.Text
p.editMessage(ctx, info.TaskID(), message, false)
}
func (p *Progress) editMessage(ctx context.Context, taskID string, message renderedBatchMessage, cancellable bool) {
if message.Err != nil {
log.FromContext(ctx).Errorf("Failed to render batch progress message: %v", message.Err)
return
text, entities := entityBuilder.Complete()
req := &tg.MessagesEditMessageRequest{
ID: p.MessageID,
}
req := buildBatchEditMessageRequest(p.MessageID, taskID, message, cancellable)
if ext := tgutil.ExtFromContext(ctx); ext != nil {
if _, err := ext.EditMessage(p.ChatID, req); err != nil {
log.FromContext(ctx).Errorf("Failed to edit batch progress message: %v", err)
}
}
}
req.SetMessage(text)
req.SetEntities(entities)
func buildBatchEditMessageRequest(messageID int, taskID string, message renderedBatchMessage, cancellable bool) *tg.MessagesEditMessageRequest {
req := &tg.MessagesEditMessageRequest{ID: messageID}
req.SetMessage(message.Text)
if len(message.Entities) > 0 {
req.SetEntities(message.Entities)
ext := tgutil.ExtFromContext(ctx)
if ext != nil {
ext.EditMessage(p.ChatID, req)
}
if cancellable {
req.SetReplyMarkup(&tg.ReplyInlineMarkup{Rows: []tg.KeyboardButtonRow{{
Buttons: []tg.KeyboardButtonClass{tgutil.BuildCancelButton(taskID)},
}}})
}
return req
}
func buildBatchProgressText(info TaskInfo, skipped []string, activeLimit int) string {
return buildBatchProgressMessage(info, skipped, activeLimit).Text
}
func buildBatchProgressMessage(info TaskInfo, skipped []string, activeLimit int) renderedBatchMessage {
items := info.Items()
completed, waiting, downloaded, failed := itemCounts(items)
downloadSpeed, uploadSpeed := aggregateSpeeds(items)
if activeLimit < 1 {
activeLimit = 1
}
total := len(items) + len(skipped)
downloadSpeedText := formatSpeed(downloadSpeed)
uploadSpeedText := formatSpeed(uploadSpeed)
header := localizedProgressMarkup(i18nk.BotMsgProgressBatchStatusHeader, map[string]any{
"Total": total,
"TotalSize": dlutil.FormatSize(info.ActualTotalSize()),
"Completed": completed,
"Downloaded": downloaded,
"Waiting": waiting,
"DownloadSpeed": downloadSpeedText,
"UploadSpeed": uploadSpeedText,
})
var markup strings.Builder
markup.WriteString(header)
visibleItems, hiddenTransfers, summarizedConfirming := visibleBatchItems(items, activeLimit)
for _, item := range visibleItems {
markup.WriteString("\n\n")
markup.WriteString(formatActiveItemMarkup(item, len(items)))
}
if hiddenTransfers > 0 {
markup.WriteString("\n\n")
markup.WriteString(localizedProgressMarkup(i18nk.BotMsgProgressBatchSummaryHiddenActive, map[string]any{"Count": hiddenTransfers}))
}
if summarizedConfirming > 0 {
markup.WriteString("\n\n")
markup.WriteString(localizedProgressMarkup(i18nk.BotMsgProgressBatchSummaryConfirming, map[string]any{"Count": summarizedConfirming}))
}
if failed > 0 {
markup.WriteString("\n")
markup.WriteString(localizedProgressMarkup(i18nk.BotMsgProgressBatchSummaryFailed, map[string]any{"Count": failed}))
}
if len(skipped) > 0 {
markup.WriteString("\n")
markup.WriteString(localizedProgressMarkup(i18nk.BotMsgProgressBatchSummarySkipped, map[string]any{"Count": len(skipped)}))
}
return completeBatchMessage(markup.String())
}
func buildBatchDoneMarkup(info TaskInfo, skipped []string, err error) string {
items := info.Items()
totalSize := info.ActualTotalSize()
if totalSize == 0 {
totalSize = info.TotalSize()
}
if err == nil {
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": completed,
"Skipped": totalSkipped,
"Size": dlutil.FormatSize(totalSize),
})
}
return localizedProgressMarkup(i18nk.BotMsgProgressBatchDone, map[string]any{
"Count": len(items),
"Size": dlutil.FormatSize(totalSize),
})
}
completed, _, _, failed := itemCounts(items)
incomplete := max(len(items)-completed-failed, 0)
if errors.Is(err, context.Canceled) {
return localizedProgressMarkup(i18nk.BotMsgProgressBatchCanceled, map[string]any{
"Total": len(items) + len(skipped),
"Completed": completed,
"Incomplete": incomplete,
"Skipped": len(skipped),
})
}
failedItems := make([]TaskItemProgress, 0, failed)
for _, item := range items {
if item.Phase == ItemPhaseFailed {
failedItems = append(failedItems, item)
}
}
if len(failedItems) > 1 && failedItems[0].FailureStage == FailureStageBatchUpload {
return localizedProgressMarkup(i18nk.BotMsgProgressBatchFailedGroup, map[string]any{
"Affected": len(failedItems),
"Reason": displayError(firstError(failedItems), err),
"Completed": completed,
"Failed": len(failedItems),
"Incomplete": incomplete,
})
}
if len(failedItems) == 0 {
return localizedProgressMarkup(i18nk.BotMsgProgressBatchFailedTask, map[string]any{
"Reason": displayError("", err),
"Completed": completed,
"Incomplete": incomplete,
})
}
item := failedItems[0]
return localizedProgressMarkup(i18nk.BotMsgProgressBatchFailedItem, map[string]any{
"Index": item.Index,
"Name": truncateFilename(item.Name, maxDisplayNameRunes),
"Stage": failureStageLabel(item.FailureStage),
"Progress": failureProgress(item),
"Speed": failureSpeed(item),
"Reason": displayError(item.Error, err),
"Completed": completed,
"Failed": failed,
"Incomplete": incomplete,
})
}
func buildBatchDoneMessage(info TaskInfo, skipped []string, err error) renderedBatchMessage {
return completeBatchMessage(buildBatchDoneMarkup(info, skipped, err))
}
func formatActiveItemMarkup(item TaskItemProgress, total int) string {
data := map[string]any{
"Index": item.Index,
"Total": total,
"Name": truncateFilename(item.Name, maxDisplayNameRunes),
"Speed": formatSpeed(itemSpeed(item)),
"Progress": itemPercent(item),
"Bar": textProgressBar(itemPercent(item)),
"Current": dlutil.FormatSize(itemBytes(item)),
"Size": dlutil.FormatSize(item.Size),
"Attempt": min(max(item.RetryAttempt, 1), max(item.RetryLimit, 1)),
"Limit": max(item.RetryLimit, 1),
"Reason": truncateRunes(item.Error, maxDisplayErrorRunes),
}
switch item.Phase {
case ItemPhaseDownloading:
if item.Size <= 0 {
return localizedProgressMarkup(i18nk.BotMsgProgressBatchItemDownloadingUnknown, data)
}
return localizedProgressMarkup(i18nk.BotMsgProgressBatchItemDownloading, data)
case ItemPhaseTransferring:
if item.Size <= 0 {
return localizedProgressMarkup(i18nk.BotMsgProgressBatchItemTransferringUnknown, data)
}
return localizedProgressMarkup(i18nk.BotMsgProgressBatchItemTransferring, data)
case ItemPhaseUploading:
return localizedProgressMarkup(i18nk.BotMsgProgressBatchItemUploading, data)
case ItemPhaseRetrying:
return localizedProgressMarkup(i18nk.BotMsgProgressBatchItemRetrying, data)
case ItemPhaseConfirming:
return localizedProgressMarkup(i18nk.BotMsgProgressBatchItemConfirming, data)
default:
return ""
}
}
func localizedProgressMarkup(key i18nk.Key, data map[string]any) string {
return i18n.T(key, tgutil.EscapeHTMLTemplateData(data))
}
func visibleBatchItems(items []TaskItemProgress, limit int) (visible []TaskItemProgress, hiddenTransfers, summarizedConfirming int) {
visible = make([]TaskItemProgress, 0, limit)
transferCount := 0
confirmingCount := 0
for _, item := range items {
switch {
case isTransferPhase(item.Phase):
transferCount++
if len(visible) < limit {
visible = append(visible, item)
}
case item.Phase == ItemPhaseConfirming:
confirmingCount++
}
}
hiddenTransfers = transferCount - len(visible)
if confirmingCount == 1 && len(visible) < limit {
for _, item := range items {
if item.Phase == ItemPhaseConfirming {
visible = append(visible, item)
return visible, hiddenTransfers, 0
}
}
}
return visible, hiddenTransfers, confirmingCount
}
func completeBatchMessage(markup string) renderedBatchMessage {
text, entities, err := tgutil.RenderHTML(markup)
return renderedBatchMessage{Text: text, Entities: entities, Err: err}
}
func itemCounts(items []TaskItemProgress) (completed, waiting, downloaded, failed int) {
for _, item := range items {
switch item.Phase {
case ItemPhaseCompleted:
completed++
case ItemPhaseWaiting:
waiting++
case ItemPhaseDownloaded:
downloaded++
case ItemPhaseFailed:
failed++
}
}
return
}
func aggregateSpeeds(items []TaskItemProgress) (download, upload float64) {
for _, item := range items {
switch item.Phase {
case ItemPhaseDownloading:
download += item.DownloadSpeed
case ItemPhaseTransferring:
download += item.DownloadSpeed
upload += item.UploadSpeed
case ItemPhaseUploading:
upload += item.UploadSpeed
}
}
return
}
func isTransferPhase(phase ItemPhase) bool {
switch phase {
case ItemPhaseDownloading, ItemPhaseTransferring, ItemPhaseUploading, ItemPhaseRetrying:
return true
default:
return false
}
}
func itemBytes(item TaskItemProgress) int64 {
switch item.Phase {
case ItemPhaseDownloading, ItemPhaseTransferring:
return item.Downloaded
default:
return item.Uploaded
}
}
func itemSpeed(item TaskItemProgress) float64 {
switch item.Phase {
case ItemPhaseDownloading, ItemPhaseTransferring:
return item.DownloadSpeed
case ItemPhaseUploading:
return item.UploadSpeed
case ItemPhaseRetrying:
return item.UploadSpeed
default:
return 0
}
}
func itemPercent(item TaskItemProgress) int {
if item.Size <= 0 {
return 0
}
return int(min(itemBytes(item), item.Size) * 100 / item.Size)
}
func textProgressBar(percent int) string {
percent = min(max(percent, 0), 100)
filled := percent * progressBarWidth / 100
return strings.Repeat("🟩", filled) + strings.Repeat("⬜️", progressBarWidth-filled)
}
func formatSpeed(speed float64) string {
if speed <= 0 {
return "0 B/s"
}
return dlutil.FormatSize(int64(speed)) + "/s"
}
func truncateFilename(name string, limit int) string {
if utf8.RuneCountInString(name) <= limit {
return name
}
ext := path.Ext(name)
if utf8.RuneCountInString(ext) >= limit-2 {
return truncateRunes(name, limit-1) + "…"
}
base := strings.TrimSuffix(name, ext)
baseLimit := limit - utf8.RuneCountInString(ext) - 1
return truncateRunes(base, baseLimit) + "…" + ext
}
func truncateRunes(value string, limit int) string {
if limit <= 0 {
return ""
}
runes := []rune(value)
if len(runes) <= limit {
return value
}
return string(runes[:limit])
}
func displayError(itemError string, fallback error) string {
if itemError == "" && fallback != nil {
itemError = compactError(fallback)
}
return truncateRunes(itemError, maxDisplayErrorRunes)
}
func firstError(items []TaskItemProgress) string {
for _, item := range items {
if item.Error != "" {
return item.Error
}
}
return ""
}
func failureStageLabel(stage FailureStage) string {
switch stage {
case FailureStageDownload:
return i18n.T(i18nk.BotMsgProgressBatchFailureStageDownload, nil)
case FailureStageCache:
return i18n.T(i18nk.BotMsgProgressBatchFailureStageCache, nil)
case FailureStageUpload:
return i18n.T(i18nk.BotMsgProgressBatchFailureStageUpload, nil)
case FailureStageConfirm:
return i18n.T(i18nk.BotMsgProgressBatchFailureStageConfirm, nil)
case FailureStageBatchUpload:
return i18n.T(i18nk.BotMsgProgressBatchFailureStageBatchUpload, nil)
default:
return i18n.T(i18nk.BotMsgProgressBatchFailureStageInternal, nil)
}
}
func failureProgress(item TaskItemProgress) string {
if item.Size <= 0 {
return dlutil.FormatSize(failureBytes(item))
}
return fmt.Sprintf("%d%%", min(failureBytes(item), item.Size)*100/item.Size)
}
func failureSpeed(item TaskItemProgress) string {
if item.FailureStage == FailureStageDownload || item.FailureStage == FailureStageCache {
return formatSpeed(item.DownloadSpeed)
}
return formatSpeed(item.UploadSpeed)
}
func failureBytes(item TaskItemProgress) int64 {
if item.FailureStage == FailureStageDownload || item.FailureStage == FailureStageCache {
return item.Downloaded
}
return item.Uploaded
}
func visibleActiveItems() int {
return min(max(config.C().Workers, 1), maxVisibleActiveItems)
}
func NewProgressTracker(messageID int, chatID int64) ProgressTracker {

View File

@@ -1,335 +0,0 @@
package batchtfile
import (
"context"
"errors"
"strings"
"sync"
"testing"
"time"
"github.com/gotd/td/tg"
"github.com/krau/SaveAny-Bot/common/i18n"
"github.com/krau/SaveAny-Bot/pkg/tfile"
)
type progressRegressionRecorder struct {
mu sync.Mutex
startTotal int64
notifications []int64
}
func (*progressRegressionRecorder) OnStart(context.Context, TaskInfo) {}
func (*progressRegressionRecorder) OnProgress(context.Context, TaskInfo) {}
func (*progressRegressionRecorder) OnDone(context.Context, TaskInfo, error) {}
func (r *progressRegressionRecorder) OnUploadStart(_ context.Context, _ TaskInfo, total int64) {
r.mu.Lock()
defer r.mu.Unlock()
r.startTotal = total
}
func (r *progressRegressionRecorder) OnUploadProgress(_ context.Context, _ TaskInfo, uploaded, _ int64) {
r.mu.Lock()
defer r.mu.Unlock()
r.notifications = append(r.notifications, uploaded)
}
type orderedProgressRegressionRecorder struct {
firstEntered chan struct{}
releaseFirst chan struct{}
secondEntered chan struct{}
mu sync.Mutex
notifications []int64
}
func (*orderedProgressRegressionRecorder) OnStart(context.Context, TaskInfo) {}
func (*orderedProgressRegressionRecorder) OnProgress(context.Context, TaskInfo) {}
func (*orderedProgressRegressionRecorder) OnDone(context.Context, TaskInfo, error) {}
func (*orderedProgressRegressionRecorder) OnUploadStart(context.Context, TaskInfo, int64) {
}
func (r *orderedProgressRegressionRecorder) OnUploadProgress(_ context.Context, _ TaskInfo, uploaded, _ int64) {
if uploaded == 100 {
close(r.firstEntered)
<-r.releaseFirst
}
if uploaded == 200 {
close(r.secondEntered)
}
r.mu.Lock()
r.notifications = append(r.notifications, uploaded)
r.mu.Unlock()
}
func TestBatchProgressShowsTransferSpeedAndSize(t *testing.T) {
useProgressRegressionLocale(t)
task := newProgressRegressionTask(nil,
progressRegressionFile{"downloading", 1000},
progressRegressionFile{"uploading", 1000},
progressRegressionFile{"waiting", 1000},
)
started := time.Unix(100, 0)
task.markItemActive("downloading", false, started)
task.recordItemDownload("downloading", 500, started.Add(time.Second))
task.recordItemDownloaded("uploading", 1000)
task.recordItemUpload("uploading", 0, 1000, started.Add(time.Second))
task.recordItemUpload("uploading", 250, 1000, started.Add(2*time.Second))
message := buildBatchProgressMessage(task, nil, 2)
if message.Err != nil {
t.Fatalf("buildBatchProgressMessage() failed: %v", message.Err)
}
assertProgressRegressionContains(t, message.Text,
"状态:✅ 0 📥 0 ⏳ 1",
"⬇️ 1/3 下载中",
"速度500 B/s",
"大小500 B / 1000 B",
"⬆️ 2/3 上传中",
"速度250 B/s",
"大小250 B / 1000 B",
)
bold, _, blockquote, _ := batchEntityCounts(message.Entities)
if bold != 3 || blockquote != 2 {
t.Fatalf("entity counts = bold:%d blockquote:%d, want bold:3 blockquote:2", bold, blockquote)
}
}
func TestBatchProgressHeaderShowsTotalSize(t *testing.T) {
useProgressRegressionLocale(t)
task := newProgressRegressionTask(nil,
progressRegressionFile{"first", 1024},
progressRegressionFile{"second", 1024},
)
message := buildBatchProgressMessage(task, nil, 2)
if message.Err != nil {
t.Fatalf("buildBatchProgressMessage() failed: %v", message.Err)
}
assertProgressRegressionContains(t, message.Text,
"文件2 总大小2.00 KB",
)
i18n.Init("en")
english := buildBatchProgressMessage(task, nil, 2)
if english.Err != nil {
t.Fatalf("English batch template failed: %v", english.Err)
}
assertProgressRegressionContains(t, english.Text,
"Files: 2 | Total size: 2.00 KB",
)
}
func TestBatchProgressLimitsRowsWithoutHidingActiveUpload(t *testing.T) {
useProgressRegressionLocale(t)
task := newProgressRegressionTask(nil,
progressRegressionFile{"confirm-01", 100},
progressRegressionFile{"confirm-02", 100},
progressRegressionFile{"uploading", 100},
progressRegressionFile{"downloading", 100},
)
started := time.Unix(100, 0)
task.recordItemUpload("confirm-01", 100, 100, started)
task.recordItemUpload("confirm-02", 100, 100, started)
task.recordItemUpload("uploading", 40, 100, started.Add(time.Second))
task.markItemActive("downloading", false, started)
message := buildBatchProgressText(task, nil, 2)
assertProgressRegressionContains(t, message,
"uploading.bin",
"downloading.bin",
"☁️ 已上传等待整组发送2",
)
if strings.Contains(message, "confirm-01.bin") || strings.Contains(message, "confirm-02.bin") {
t.Fatalf("confirmation rows displaced active transfers:\n%s", message)
}
}
func TestBatchProgressTemplateOwnsStylesAndEscapesValues(t *testing.T) {
useProgressRegressionLocale(t)
fileID := `<b>A&B</b>`
task := newProgressRegressionTask(nil, progressRegressionFile{fileID, 100})
task.markItemRetry(fileID, FailureStageUpload, 1, 3, errors.New(`<i>remote & failed</i>`))
message := buildBatchProgressMessage(task, nil, 1)
if message.Err != nil {
t.Fatalf("buildBatchProgressMessage() failed: %v", message.Err)
}
assertProgressRegressionContains(t, message.Text,
`<b>A&B</b>.bin`,
`<i>remote & failed</i>`,
)
bold, _, blockquote, italic := batchEntityCounts(message.Entities)
if bold != 2 || blockquote != 1 || italic != 0 {
t.Fatalf("entity counts = bold:%d blockquote:%d italic:%d", bold, blockquote, italic)
}
i18n.Init("en")
english := buildBatchProgressMessage(task, nil, 1)
if english.Err != nil {
t.Fatalf("English batch template failed: %v", english.Err)
}
assertProgressRegressionContains(t, english.Text, "📦 Processing", "Retrying upload", `<b>A&B</b>.bin`)
}
func TestDownloadProgressContinuesAfterUploadStarts(t *testing.T) {
useProgressRegressionLocale(t)
progress := new(Progress)
task := newProgressRegressionTask(progress,
progressRegressionFile{"uploading", 100},
progressRegressionFile{"downloading", 100},
)
progress.OnStart(t.Context(), task)
task.recordDownloadComplete("uploading", 100)
task.uploadCallback(t.Context(), "uploading")(50, 100)
started := time.Unix(100, 0)
task.markItemActive("downloading", false, started)
task.recordItemDownload("downloading", 50, started.Add(time.Second))
progress.updateMu.Lock()
progress.lastUpdateAt = time.Now().Add(-progressRenderInterval)
progress.updateMu.Unlock()
progress.OnProgress(t.Context(), task)
progress.updateMu.Lock()
text := progress.lastText
progress.updateMu.Unlock()
assertProgressRegressionContains(t, text,
"uploading.bin",
"🟩🟩🟩🟩🟩⬜️⬜️⬜️⬜️⬜️ 50%",
"总速度:⬇️ 50 B/s ⬆️ 0 B/s",
"🔄 另有 1 个文件正在处理",
)
}
func TestBatchUploadIgnoresOutOfOrderBytesAndAllowsRetryReset(t *testing.T) {
recorder := new(progressRegressionRecorder)
task := newProgressRegressionTask(recorder, progressRegressionFile{"file", 100})
task.recordDownloadComplete("file", 100)
callback := task.uploadCallback(t.Context(), "file")
callback(80, 100)
callback(10, 100)
if got := task.Items()[0].Uploaded; got != 80 {
t.Fatalf("out-of-order callback regressed item to %d, want 80", got)
}
task.markItemRetry("file", FailureStageUpload, 1, 3, context.DeadlineExceeded)
callback(0, 100)
callback(10, 100)
if got := task.Items()[0].Uploaded; got != 10 {
t.Fatalf("retry did not reset item progress: got %d, want 10", got)
}
recorder.mu.Lock()
defer recorder.mu.Unlock()
for index := 1; index < len(recorder.notifications); index++ {
if recorder.notifications[index] < recorder.notifications[index-1] {
t.Fatalf("aggregate progress regressed: %v", recorder.notifications)
}
}
}
func TestUploadProgressNotificationsRemainOrdered(t *testing.T) {
recorder := &orderedProgressRegressionRecorder{
firstEntered: make(chan struct{}),
releaseFirst: make(chan struct{}),
secondEntered: make(chan struct{}),
}
task := newProgressRegressionTask(recorder,
progressRegressionFile{"first", 100},
progressRegressionFile{"second", 100},
)
task.recordDownloadComplete("first", 100)
task.recordDownloadComplete("second", 100)
first := task.uploadCallback(t.Context(), "first")
second := task.uploadCallback(t.Context(), "second")
var wait sync.WaitGroup
wait.Go(func() {
first(100, 100)
})
<-recorder.firstEntered
wait.Go(func() {
second(100, 100)
})
overtook := false
select {
case <-recorder.secondEntered:
overtook = true
case <-time.After(100 * time.Millisecond):
}
close(recorder.releaseFirst)
wait.Wait()
if overtook {
t.Fatal("later aggregate notification overtook the first callback")
}
recorder.mu.Lock()
defer recorder.mu.Unlock()
if got := recorder.notifications; len(got) != 2 || got[0] != 100 || got[1] != 200 {
t.Fatalf("upload notifications = %v, want [100 200]", got)
}
}
func TestBatchUploadUsesActualSizeWhenMetadataIsUnknown(t *testing.T) {
recorder := new(progressRegressionRecorder)
task := newProgressRegressionTask(recorder, progressRegressionFile{"photo", 0})
task.recordDownloadComplete("photo", 25)
task.uploadCallback(t.Context(), "photo")(25, 25)
recorder.mu.Lock()
defer recorder.mu.Unlock()
if recorder.startTotal != 25 {
t.Fatalf("upload start total = %d, want actual size 25", recorder.startTotal)
}
if got := task.ActualTotalSize(); got != 25 {
t.Fatalf("actual total size = %d, want 25", got)
}
}
type progressRegressionFile struct {
id string
size int64
}
func newProgressRegressionTask(progress ProgressTracker, files ...progressRegressionFile) *Task {
elems := make([]TaskElement, 0, len(files))
for _, file := range files {
elems = append(elems, TaskElement{
ID: file.id,
File: tfile.NewTGFile(nil, nil, file.size, file.id+".bin"),
})
}
return NewBatchTGFileTask("progress-regression", context.Background(), elems, progress, true)
}
func useProgressRegressionLocale(t *testing.T) {
t.Helper()
i18n.Init("zh-Hans")
t.Cleanup(func() { i18n.Init("zh-Hans") })
}
func assertProgressRegressionContains(t *testing.T, value string, wants ...string) {
t.Helper()
for _, want := range wants {
if !strings.Contains(value, want) {
t.Fatalf("text does not contain %q:\n%s", want, value)
}
}
}
func batchEntityCounts(entities []tg.MessageEntityClass) (bold, code, blockquote, italic int) {
for _, messageEntity := range entities {
switch messageEntity.(type) {
case *tg.MessageEntityBold:
bold++
case *tg.MessageEntityCode:
code++
case *tg.MessageEntityBlockquote:
blockquote++
case *tg.MessageEntityItalic:
italic++
}
}
return
}

View File

@@ -2,14 +2,11 @@ package batchtfile
import (
"context"
"encoding/json"
"fmt"
"path/filepath"
"sync"
"sync/atomic"
"github.com/charmbracelet/log"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/core"
"github.com/krau/SaveAny-Bot/pkg/enums/tasktype"
@@ -21,35 +18,25 @@ import (
var _ core.Executable = (*Task)(nil)
type TaskElement struct {
ID string
Storage storage.Storage
Path string
File tfile.TGFile
localPath string
stream bool
sourceGroupKey string
sourceCaption string
preserveCaption bool
ID string
Storage storage.Storage
Path string
File tfile.TGFile
localPath string
stream bool
}
type Task struct {
ID string
ctx context.Context
elems []TaskElement
Progress ProgressTracker
IgnoreErrors bool // if true, errors during processing will be ignored
downloaded atomic.Int64
totalSize int64
uploadTotalSize atomic.Int64
processing map[string]TaskElementInfo
processingMu sync.RWMutex
itemStates []itemProgressState
itemIndex map[string]int
itemMu sync.RWMutex
uploadOnce sync.Once
uploadMu sync.Mutex
uploaded map[string]int64
overwrite bool // recovered: overwrite storage targets instead of uniquifying
ID string
ctx context.Context
elems []TaskElement
Progress ProgressTracker
IgnoreErrors bool // if true, errors during processing will be ignored
downloaded atomic.Int64
totalSize int64
processing map[string]TaskElementInfo
processingMu sync.RWMutex
failed map[string]error // [TODO] errors for each element
}
// Title implements core.Exectable.
@@ -61,48 +48,12 @@ func (t *Task) Type() tasktype.TaskType {
return tasktype.TaskTypeTgfiles
}
// completedElementIDs returns the element IDs whose upload finished, for
// persisting upload progress so recovery can skip them.
func (t *Task) completedElementIDs() []string {
t.itemMu.RLock()
defer t.itemMu.RUnlock()
var ids []string
for _, item := range t.itemStates {
if item.phase == ItemPhaseCompleted {
ids = append(ids, item.id)
}
}
return ids
}
// persistElementDone records an element's completed upload in the persisted
// payload so a restart does not re-upload it.
func (t *Task) persistElementDone(ctx context.Context, elemID string) {
err := core.UpdateTaskPayload(ctx, t.ID, func(payload []byte) ([]byte, error) {
var p taskPayload
if err := json.Unmarshal(payload, &p); err != nil {
return nil, err
}
for _, id := range p.Done {
if id == elemID {
return payload, nil
}
}
p.Done = append(p.Done, elemID)
return json.Marshal(p)
})
if err != nil {
log.FromContext(ctx).Warnf("Failed to persist element completion %s: %v", elemID, err)
}
}
func NewTaskElement(
stor storage.Storage,
path string,
file tfile.TGFile,
) (*TaskElement, error) {
id := xid.New().String()
groupKey, caption, preserveCaption := sourceMetadata(file)
_, ok := stor.(storage.StorageCannotStream)
if !config.C().Stream || ok {
cachePath, err := filepath.Abs(filepath.Join(config.C().Temp.BasePath, fmt.Sprintf("%s_%s", id, file.Name())))
@@ -110,42 +61,22 @@ func NewTaskElement(
return nil, fmt.Errorf("failed to get absolute path for cache: %w", err)
}
return &TaskElement{
ID: id,
Storage: stor,
Path: path,
File: file,
localPath: cachePath,
sourceGroupKey: groupKey,
sourceCaption: caption,
preserveCaption: preserveCaption,
ID: id,
Storage: stor,
Path: path,
File: file,
localPath: cachePath,
}, nil
}
return &TaskElement{
ID: id,
Storage: stor,
Path: path,
File: file,
stream: true,
sourceGroupKey: groupKey,
sourceCaption: caption,
preserveCaption: preserveCaption,
ID: id,
Storage: stor,
Path: path,
File: file,
stream: true,
}, nil
}
func sourceMetadata(file tfile.TGFile) (groupKey, caption string, preserveCaption bool) {
messageFile, ok := file.(tfile.TGFileMessage)
if !ok || messageFile.Message() == nil {
return "", "", false
}
msg := messageFile.Message()
groupID, grouped := msg.GetGroupedID()
if !grouped || groupID == 0 {
return "", "", false
}
chatID := tgutil.ChatIdFromPeer(msg.GetPeerID())
return fmt.Sprintf("%T:%d:%d", msg.GetPeerID(), chatID, groupID), msg.GetMessage(), true
}
func NewBatchTGFileTask(
id string,
ctx context.Context,
@@ -153,7 +84,6 @@ func NewBatchTGFileTask(
progress ProgressTracker,
ignoreErrors bool,
) *Task {
itemStates, itemIndex := newItemProgressStates(files)
task := &Task{
ID: id,
ctx: ctx,
@@ -168,11 +98,9 @@ func NewBatchTGFileTask(
return total
}(),
processing: make(map[string]TaskElementInfo),
itemStates: itemStates,
itemIndex: itemIndex,
uploaded: make(map[string]int64),
IgnoreErrors: ignoreErrors,
processingMu: sync.RWMutex{},
failed: make(map[string]error),
}
return task
}

View File

@@ -27,10 +27,8 @@ type TaskInfo interface {
TaskID() string
TotalSize() int64
Downloaded() int64
ActualTotalSize() int64
Count() int
Processing() []TaskElementInfo
Items() []TaskItemProgress
}
func (t *Task) TaskID() string {

View File

@@ -1,70 +0,0 @@
package batchtfile
import (
"context"
"time"
)
// UploadProgressTracker optionally extends a batch progress tracker with a
// distinct aggregate upload phase.
type UploadProgressTracker interface {
OnUploadStart(ctx context.Context, info TaskInfo, total int64)
OnUploadProgress(ctx context.Context, info TaskInfo, uploaded, total int64)
}
func (t *Task) startUpload(ctx context.Context) {
tracker, ok := t.Progress.(UploadProgressTracker)
if !ok {
return
}
t.uploadMu.Lock()
defer t.uploadMu.Unlock()
t.uploadOnce.Do(func() {
tracker.OnUploadStart(ctx, t, t.uploadTotalSize.Load())
})
}
func (t *Task) uploadCallback(ctx context.Context, id string) func(uploaded, total int64) {
return func(uploaded, total int64) {
tracker, ok := t.Progress.(UploadProgressTracker)
if !ok || uploaded < 0 {
return
}
t.startUpload(ctx)
if total > 0 && uploaded > total {
uploaded = total
}
t.uploadMu.Lock()
defer t.uploadMu.Unlock()
becameConfirming := t.recordItemUpload(id, uploaded, total, time.Now())
if t.uploaded == nil {
t.uploaded = make(map[string]int64)
}
previous, tracked := t.uploaded[id]
if !tracked || uploaded > previous {
t.uploaded[id] = uploaded
}
var aggregate int64
for _, current := range t.uploaded {
aggregate += current
}
uploadTotal := t.uploadTotalSize.Load()
if aggregate > uploadTotal {
aggregate = uploadTotal
}
tracker.OnUploadProgress(ctx, t, aggregate, uploadTotal)
if becameConfirming {
t.notifyStateChange(ctx)
}
}
}
func (t *Task) recordDownloadComplete(id string, uploadSize int64) {
t.uploadMu.Lock()
defer t.uploadMu.Unlock()
if uploadSize > 0 {
t.uploadTotalSize.Add(uploadSize)
}
t.recordItemDownloaded(id, uploadSize)
}

View File

@@ -0,0 +1,32 @@
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
}

View File

@@ -76,11 +76,12 @@ func (t *Task) Execute(ctx context.Context) error {
eg.SetLimit(config.C().Workers)
for _, file := range t.files {
eg.Go(func() error {
t.processingMu.Lock()
t.processingMu.RLock()
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() {
@@ -89,6 +90,7 @@ 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
@@ -97,7 +99,6 @@ 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
})
}

View File

@@ -15,7 +15,6 @@ 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"
)
@@ -103,7 +102,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 !progressutil.ShouldUpdate(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
if !shouldUpdateProgress(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
return
}
percent := int((info.DownloadedBytes() * 100) / info.TotalBytes())
@@ -116,10 +115,7 @@ 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(i18n.T(i18nk.BotMsgProgressSizeWithFiles, map[string]any{
"Size": fmt.Sprintf("%.2f MB", float64(info.TotalBytes())/(1024*1024)),
"Count": info.TotalFiles(),
})),
styling.Code(fmt.Sprintf("%.2f MB (%d个文件)", float64(info.TotalBytes())/(1024*1024), info.TotalFiles())),
styling.Plain(i18n.T(i18nk.BotMsgProgressProcessingListPrefix, nil)),
func() styling.StyledTextOption {
var lines []string

View File

@@ -45,6 +45,7 @@ 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.
@@ -126,6 +127,7 @@ func NewTask(
client: http.DefaultClient,
processing: make(map[string]*File),
processingMu: sync.RWMutex{},
failed: make(map[string]error),
totalFiles: int64(len(files)),
}
}

View File

@@ -207,3 +207,34 @@ 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
}

View File

@@ -30,20 +30,21 @@ func (t *Task) Execute(ctx context.Context) error {
eg.SetLimit(config.C().Workers)
for _, resource := range t.item.Resources {
eg.Go(func() error {
resourceID := resource.ID()
t.processingMu.Lock()
if t.processing[resourceID] != nil {
t.processingMu.Unlock()
return fmt.Errorf("resource %s is already being processed", resourceID)
t.processingMu.RLock()
if t.processing[resource.ID()] != nil {
return fmt.Errorf("resource %s is already being processed", resource.ID())
}
t.processing[resourceID] = &resource
t.processingMu.RUnlock()
t.processingMu.Lock()
t.processing[resource.ID()] = &resource
t.processingMu.Unlock()
defer func() {
t.processingMu.Lock()
delete(t.processing, resourceID)
delete(t.processing, resource.URL)
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
@@ -52,7 +53,6 @@ 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
})
}

View File

@@ -15,10 +15,40 @@ 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)
@@ -43,10 +73,7 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
styling.Plain(i18n.T(i18nk.BotMsgProgressParsedStartPrefix, map[string]any{
"Site": info.Site(),
})),
styling.Code(i18n.T(i18nk.BotMsgProgressSizeWithResources, map[string]any{
"Size": fmt.Sprintf("%.2f MB", float64(info.TotalBytes())/(1024*1024)),
"Count": info.TotalResources(),
})),
styling.Code(fmt.Sprintf("%.2f MB (%d个资源)", float64(info.TotalBytes())/(1024*1024), info.TotalResources())),
); err != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", err)
return
@@ -74,7 +101,7 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
if !progressutil.ShouldUpdate(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
if !shouldUpdateProgress(info.TotalBytes(), info.DownloadedBytes(), int(p.lastUpdatePercent.Load())) {
return
}
percent := int((info.DownloadedBytes() * 100) / info.TotalBytes())
@@ -87,10 +114,7 @@ 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(i18n.T(i18nk.BotMsgProgressSizeWithResources, map[string]any{
"Size": fmt.Sprintf("%.2f MB", float64(info.TotalBytes())/(1024*1024)),
"Count": info.TotalResources(),
})),
styling.Code(fmt.Sprintf("%.2f MB (%d个文件)", float64(info.TotalBytes())/(1024*1024), info.TotalResources())),
styling.Plain(i18n.T(i18nk.BotMsgProgressProcessingListPrefix, nil)),
func() styling.StyledTextOption {
var lines []string

View File

@@ -33,6 +33,7 @@ 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.
@@ -83,5 +84,6 @@ func NewTask(
progress: progressTracker,
processing: make(map[string]ResourceInfo),
processingMu: sync.RWMutex{},
failed: make(map[string]error),
}
}

View File

@@ -18,9 +18,7 @@ import (
func (t *Task) Execute(ctx context.Context) error {
logger := log.FromContext(ctx)
logger.Infof("Starting Telegraph task %s", t.PhPath)
if t.progress != nil {
t.progress.OnStart(ctx, t)
}
t.progress.OnStart(ctx, t)
eg, gctx := errgroup.WithContext(ctx)
eg.SetLimit(config.C().Workers)
for i, pic := range t.Pics {
@@ -31,9 +29,7 @@ func (t *Task) Execute(ctx context.Context) error {
return fmt.Errorf("failed to process picture %s: %w", pic, err)
}
downloaded := t.downloaded.Add(1)
if t.progress != nil {
t.progress.OnProgress(gctx, t)
}
t.progress.OnProgress(gctx, t)
taskevent.Emit(gctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,
@@ -49,9 +45,7 @@ func (t *Task) Execute(ctx context.Context) error {
} else {
logger.Infof("Telegraph task %s completed successfully", t.PhPath)
}
if t.progress != nil {
t.progress.OnDone(ctx, t, err)
}
t.progress.OnDone(ctx, t, err)
return err
}

View File

@@ -1,31 +0,0 @@
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)
}
}

View File

@@ -11,7 +11,6 @@ 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"
)
@@ -61,7 +60,7 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
if !progressutil.ShouldUpdateCount(info.Downloaded(), int64(info.TotalPics())) {
if !shouldUpdateProgress(info.Downloaded(), int64(info.TotalPics())) {
return
}
log.FromContext(ctx).Debugf("Progress update: %s, %d/%d", info.TaskID(), info.Downloaded(), info.TotalPics())

View File

@@ -0,0 +1,13 @@
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
}

View File

@@ -1,42 +0,0 @@
package tfile
import (
"testing"
"github.com/gotd/td/tg"
tfilepkg "github.com/krau/SaveAny-Bot/pkg/tfile"
)
func TestSourceCaption(t *testing.T) {
tests := []struct {
name string
file tfilepkg.TGFile
want string
ok bool
}{
{
name: "original caption",
file: tfilepkg.NewTGFile(nil, nil, 0, "video.mov", tfilepkg.WithMessage(&tg.Message{Message: "original caption"})),
want: "original caption",
ok: true,
},
{
name: "empty caption suppresses storage fallback",
file: tfilepkg.NewTGFile(nil, nil, 0, "video.mov", tfilepkg.WithMessage(&tg.Message{})),
ok: true,
},
{
name: "file without source message",
file: tfilepkg.NewTGFile(nil, nil, 0, "video.mov"),
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, ok := sourceCaption(tt.file)
if ok != tt.ok || got != tt.want {
t.Fatalf("sourceCaption() = (%q, %v), want (%q, %v)", got, ok, tt.want, tt.ok)
}
})
}
}

View File

@@ -1,97 +0,0 @@
package tfile
import (
"context"
"encoding/json"
"fmt"
"path/filepath"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/core"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
tfilepkg "github.com/krau/SaveAny-Bot/pkg/tfile"
"github.com/krau/SaveAny-Bot/storage"
)
type taskPayload struct {
Kind string `json:"kind"` // "file"
ID string `json:"id"`
Storage string `json:"storage"`
Path string `json:"path"`
File tfilepkg.FilePayload `json:"file"`
ChatID int64 `json:"chat_id"`
MessageID int `json:"message_id"`
Overwrite bool `json:"overwrite"`
Caption string `json:"caption"`
}
type taskCodec struct{}
// TaskCodec serializes single-file tasks. It is registered together with the
// batch codec under TaskTypeTgfiles (see core/tasks/batchtfile/codec.go).
var TaskCodec core.TaskCodec = taskCodec{}
func (taskCodec) Marshal(task core.Executable) ([]byte, error) {
t, ok := task.(*Task)
if !ok {
return nil, fmt.Errorf("unexpected task type %T", task)
}
filePayload, ok := tfilepkg.FilePayloadOf(t.File)
if !ok {
return nil, fmt.Errorf("file %T is not serializable", t.File)
}
p := taskPayload{
Kind: "file",
ID: t.ID,
Storage: t.Storage.Name(),
Path: t.Path,
File: filePayload,
}
if overwrite, ok := t.Ctx.Value(ctxkey.OverwriteExisting).(bool); ok {
p.Overwrite = overwrite
}
if caption, ok := sourceCaption(t.File); ok {
p.Caption = caption
}
if progress, ok := t.Progress.(*Progress); ok {
p.ChatID = progress.ChatID
p.MessageID = progress.MessageID
}
return json.Marshal(p)
}
func (taskCodec) Unmarshal(data []byte) (core.Executable, error) {
var p taskPayload
if err := json.Unmarshal(data, &p); err != nil {
return nil, fmt.Errorf("invalid task payload: %w", err)
}
dler := core.DownloaderClient()
if dler == nil {
return nil, fmt.Errorf("no downloader client available")
}
file := tfilepkg.FileFromPayload(p.File, dler)
stor, err := storage.GetStorageByName(context.Background(), p.Storage)
if err != nil {
return nil, fmt.Errorf("storage %q: %w", p.Storage, err)
}
var progress ProgressTracker
if p.ChatID != 0 {
progress = NewProgressTrack(p.MessageID, p.ChatID)
}
localPath, err := filepath.Abs(filepath.Join(config.C().Temp.BasePath, fmt.Sprintf("%s_%s", p.ID, file.Name())))
if err != nil {
return nil, fmt.Errorf("failed to build cache path: %w", err)
}
return &Task{
ID: p.ID,
Ctx: context.Background(),
File: file,
Storage: stor,
Path: p.Path,
Progress: progress,
stream: false, // recovered tasks always download to cache first
localPath: localPath,
overwrite: p.Overwrite,
caption: p.Caption,
}, nil
}

View File

@@ -1,75 +0,0 @@
package tfile
import (
"context"
"fmt"
"os"
"github.com/charmbracelet/log"
"github.com/krau/SaveAny-Bot/common/tdler"
"github.com/krau/SaveAny-Bot/common/utils/dlutil"
"github.com/krau/SaveAny-Bot/common/utils/fsutil"
"github.com/krau/SaveAny-Bot/config"
)
// download fetches the file into the cache path. It resumes from a partial
// .part download tracked by a resume bitmap, and reuses a complete cache
// file (e.g. when the previous run was interrupted during upload).
func (t *Task) download(ctx context.Context) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", t.File.Name()))
if t.File.Size() > 0 {
if stat, err := os.Stat(t.localPath); err == nil && stat.Size() == t.File.Size() {
logger.Info("Cache file already complete, skipping download")
return nil
}
}
if t.File.Size() <= 0 {
// Unknown size (e.g. photos) cannot be resumed; use the plain downloader.
localFile, err := fsutil.CreateFile(t.localPath)
if err != nil {
return fmt.Errorf("failed to create local file: %w", err)
}
defer localFile.Close()
wrAt := newWriterAt(ctx, localFile, t.Progress, t)
if _, err := tdler.NewDownloader(t.File).Parallel(ctx, wrAt); err != nil {
return err
}
logger.Info("File downloaded successfully")
return nil
}
partPath := t.localPath + ".part"
// 不截断已存在的 .part: 位图标记的已完成块依赖既有字节。
localFile, err := os.OpenFile(partPath, os.O_CREATE|os.O_RDWR, 0o644)
if err != nil {
return fmt.Errorf("failed to create local file: %w", err)
}
wrAt := newWriterAt(ctx, localFile, t.Progress, t)
err = tdler.DownloadResumable(
ctx, t.File, wrAt,
dlutil.BestThreads(t.File.Size(), config.C().Threads),
tdler.ResumeStatePath(partPath),
)
closeErr := localFile.Close()
if err != nil {
return err
}
if closeErr != nil {
return fmt.Errorf("failed to close cache file: %w", closeErr)
}
stat, err := os.Stat(partPath)
if err != nil {
return fmt.Errorf("failed to stat downloaded file: %w", err)
}
if stat.Size() != t.File.Size() {
return fmt.Errorf("downloaded size %d does not match expected %d", stat.Size(), t.File.Size())
}
if err := os.Rename(partPath, t.localPath); err != nil {
return fmt.Errorf("failed to finalize download: %w", err)
}
// 清理位图是尽力而为: 下载已完成, 清理失败不应使任务失败。
if err := tdler.RemoveResumeState(tdler.ResumeStatePath(partPath)); err != nil {
logger.Warnf("Failed to remove resume state: %v", err)
}
logger.Info("File downloaded successfully")
return nil
}

View File

@@ -3,31 +3,19 @@ package tfile
import (
"context"
"fmt"
"io"
"os"
"path"
"github.com/charmbracelet/log"
"github.com/duke-git/lancet/v2/retry"
"github.com/krau/SaveAny-Bot/common/tdler"
"github.com/krau/SaveAny-Bot/common/utils/fsutil"
"github.com/krau/SaveAny-Bot/common/utils/ioutil"
"github.com/krau/SaveAny-Bot/config"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
"github.com/krau/SaveAny-Bot/pkg/storagetypes"
tfilepkg "github.com/krau/SaveAny-Bot/pkg/tfile"
"github.com/krau/SaveAny-Bot/storage"
)
func (t *Task) Execute(ctx context.Context) (err error) {
func (t *Task) Execute(ctx context.Context) error {
logger := log.FromContext(ctx).WithPrefix(fmt.Sprintf("file[%s]", t.File.Name()))
defer func() {
if t.Progress != nil {
t.Progress.OnDone(ctx, t, err)
}
}()
if t.overwrite {
ctx = storage.WithOverwrite(ctx)
}
if t.Progress != nil {
t.Progress.OnStart(ctx, t)
}
@@ -36,50 +24,46 @@ func (t *Task) Execute(ctx context.Context) (err error) {
}
logger.Info("Starting file download")
if err := t.download(ctx); err != nil {
localFile, err := fsutil.CreateFile(t.localPath)
if err != nil {
return fmt.Errorf("failed to create local file: %w", err)
}
defer func() {
if err := localFile.CloseAndRemove(); err != nil {
logger.Errorf("Failed to close local file: %v", err)
}
}()
wrAt := newWriterAt(ctx, localFile, t.Progress, t)
defer func() {
if t.Progress != nil {
t.Progress.OnDone(ctx, t, err)
}
}()
_, err = tdler.NewDownloader(t.File).Parallel(ctx, wrAt)
if err != nil {
return fmt.Errorf("failed to download file: %w", err)
}
logger.Infof("File downloaded successfully")
if path.Ext(t.File.Name()) == "" {
ext := fsutil.DetectFileExt(t.localPath)
if ext != "" {
t.Path = t.Path + ext
}
}
fileStat, err := os.Stat(t.localPath)
var fileStat os.FileInfo
fileStat, err = os.Stat(t.localPath)
if err != nil {
return fmt.Errorf("failed to get file stat: %w", err)
}
vctx := context.WithValue(ctx, ctxkey.ContentLength, fileStat.Size())
if t.caption != "" {
vctx = storagetypes.WithSourceCaption(vctx, t.caption)
} else if caption, ok := sourceCaption(t.File); ok {
vctx = storagetypes.WithSourceCaption(vctx, caption)
}
err = retry.Retry(func() error {
file, err := os.Open(t.localPath)
if err != nil {
return fmt.Errorf("failed to open cache file: %w", err)
}
defer file.Close()
uploadProgress, tracksUpload := t.Progress.(UploadProgressTracker)
if !tracksUpload {
if err = t.Storage.Save(vctx, file, t.Path); err != nil {
return fmt.Errorf("failed to save file: %w", err)
}
return nil
}
uploadProgress.OnUploadStart(vctx, t, fileStat.Size())
onProgress := func(uploaded, total int64) {
uploadProgress.OnUploadProgress(vctx, t, uploaded, total)
}
if progressSaver, ok := t.Storage.(storage.StorageProgressSaver); ok {
err = progressSaver.SaveWithProgress(vctx, file, t.Path, onProgress)
} else {
var reader io.Reader = ioutil.NewProgressReader(file, fileStat.Size(), onProgress)
err = t.Storage.Save(vctx, reader, t.Path)
}
if err != nil {
if err = t.Storage.Save(vctx, file, t.Path); err != nil {
return fmt.Errorf("failed to save file: %w", err)
}
return nil
@@ -87,17 +71,5 @@ func (t *Task) Execute(ctx context.Context) (err error) {
if err != nil {
return fmt.Errorf("failed to save file after retries: %w", err)
}
// Cache file is kept on failure so a later restart can resume upload.
if err := os.Remove(t.localPath); err != nil {
logger.Errorf("Failed to remove cache file: %v", err)
}
return nil
}
func sourceCaption(file tfilepkg.TGFile) (string, bool) {
messageFile, ok := file.(tfilepkg.TGFileMessage)
if !ok || messageFile.Message() == nil {
return "", false
}
return messageFile.Message().GetMessage(), true
}

View File

@@ -4,17 +4,16 @@ import (
"context"
"errors"
"fmt"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/charmbracelet/log"
"github.com/gotd/td/telegram/message/entity"
"github.com/gotd/td/telegram/message/styling"
"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/dlutil"
"github.com/krau/SaveAny-Bot/common/utils/progressutil"
"github.com/krau/SaveAny-Bot/common/utils/tgutil"
)
@@ -24,319 +23,153 @@ type ProgressTracker interface {
OnDone(ctx context.Context, info TaskInfo, err error)
}
// UploadProgressTracker optionally extends a task progress tracker with a
// distinct upload phase. Keeping it separate preserves compatibility with
// custom download-only trackers.
type UploadProgressTracker interface {
OnUploadStart(ctx context.Context, info TaskInfo, total int64)
OnUploadProgress(ctx context.Context, info TaskInfo, uploaded, total int64)
}
type Progress struct {
MessageID int
ChatID int64
start time.Time
lastUpdatePercent atomic.Int32
lastUpdateAt atomic.Int64
updateMu sync.Mutex
uploadAttempt int
uploadedBytes int64
actualSize int64
hasActualSize bool
}
const (
uploadProgressMinInterval = time.Second
uploadProgressMaxInterval = 3 * time.Second
singleProgressBarWidth = 10
maxSingleErrorRunes = 240
)
type singleProgressPhase int
const (
singlePhaseDownloading singleProgressPhase = iota
singlePhaseUploading
singlePhaseRetrying
)
type renderedSingleMessage struct {
Text string
Entities []tg.MessageEntityClass
Err error
}
func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
p.updateMu.Lock()
defer p.updateMu.Unlock()
p.start = time.Now()
p.lastUpdatePercent.Store(0)
p.lastUpdateAt.Store(0)
p.uploadAttempt = 0
p.uploadedBytes = 0
p.actualSize = 0
p.hasActualSize = false
log.FromContext(ctx).Debugf("Progress tracking started for message %d in chat %d", p.MessageID, p.ChatID)
p.editMessage(ctx, info.TaskID(), buildSingleProgressMessage(info, singlePhaseDownloading, 0, info.FileSize(), 0, 0), true)
entityBuilder := entity.Builder{}
var entities []tg.MessageEntityClass
if err := styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressFileStartPrefix, nil)),
styling.Code(info.FileName()),
styling.Plain(i18n.T(i18nk.BotMsgProgressSavePathPrefix, nil)),
styling.Code(fmt.Sprintf("[%s]:%s", info.StorageName(), info.StoragePath())),
styling.Plain(i18n.T(i18nk.BotMsgProgressFileSizePrefix, nil)),
styling.Code(fmt.Sprintf("%.2f MB", float64(info.FileSize())/(1024*1024))),
); err != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", err)
return
}
text, entities := entityBuilder.Complete()
req := &tg.MessagesEditMessageRequest{
ID: p.MessageID,
}
req.SetMessage(text)
req.SetEntities(entities)
req.SetReplyMarkup(&tg.ReplyInlineMarkup{
Rows: []tg.KeyboardButtonRow{
{
Buttons: []tg.KeyboardButtonClass{
tgutil.BuildCancelButton(info.TaskID()),
},
},
}},
)
ext := tgutil.ExtFromContext(ctx)
if ext != nil {
ext.EditMessage(p.ChatID, req)
return
}
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo, downloaded, total int64) {
p.updateMu.Lock()
defer p.updateMu.Unlock()
now := time.Now()
elapsed := uploadProgressMaxInterval
if lastUpdateAt := p.lastUpdateAt.Load(); lastUpdateAt > 0 {
elapsed = now.Sub(time.Unix(0, lastUpdateAt))
}
if !shouldUpdateSingleDownloadProgress(total, downloaded, int(p.lastUpdatePercent.Load()), elapsed) {
if !shouldUpdateProgress(total, downloaded, int(p.lastUpdatePercent.Load())) {
return
}
if total > 0 {
percent := int32((downloaded * 100) / total)
if p.lastUpdatePercent.Load() == percent {
return
}
p.lastUpdatePercent.Store(percent)
}
p.lastUpdateAt.Store(now.UnixNano())
log.FromContext(ctx).Debugf("Progress update: %s, %d/%d", info.FileName(), downloaded, total)
p.editMessage(ctx, info.TaskID(), buildSingleProgressMessage(
info,
singlePhaseDownloading,
downloaded,
total,
dlutil.GetSpeed(downloaded, p.start),
0,
), true)
}
func shouldUpdateSingleDownloadProgress(total, downloaded int64, lastPercent int, elapsed time.Duration) bool {
if total > 0 {
return progressutil.ShouldUpdate(total, downloaded, lastPercent)
}
return downloaded > 0 && elapsed >= uploadProgressMaxInterval
}
func (p *Progress) OnUploadStart(ctx context.Context, info TaskInfo, total int64) {
p.updateMu.Lock()
defer p.updateMu.Unlock()
p.start = time.Now()
p.lastUpdatePercent.Store(0)
p.lastUpdateAt.Store(p.start.UnixNano())
p.uploadAttempt++
p.uploadedBytes = 0
p.actualSize = max(total, 0)
p.hasActualSize = true
log.FromContext(ctx).Debugf("Upload progress tracking started: %s", info.FileName())
phase := singleUploadPhase(p.uploadAttempt)
p.editMessage(ctx, info.TaskID(), buildSingleProgressMessage(info, phase, 0, total, 0, p.uploadAttempt), true)
}
func (p *Progress) OnUploadProgress(ctx context.Context, info TaskInfo, uploaded, total int64) {
if total <= 0 || uploaded <= 0 {
percent := int32((downloaded * 100) / total)
if p.lastUpdatePercent.Load() == percent {
return
}
p.updateMu.Lock()
defer p.updateMu.Unlock()
if uploaded > total {
uploaded = total
}
if uploaded < p.uploadedBytes {
return
}
p.uploadedBytes = uploaded
now := time.Now()
lastUpdateAt := time.Unix(0, p.lastUpdateAt.Load())
lastPercent := int(p.lastUpdatePercent.Load())
if !shouldUpdateUploadProgress(total, uploaded, lastPercent, now.Sub(lastUpdateAt)) {
return
}
percent := int32((uploaded * 100) / total)
p.lastUpdatePercent.Store(percent)
p.lastUpdateAt.Store(now.UnixNano())
log.FromContext(ctx).Debugf("Upload progress update: %s, %d/%d", info.FileName(), uploaded, total)
p.editMessage(ctx, info.TaskID(), buildSingleProgressMessage(
info,
singleUploadPhase(p.uploadAttempt),
uploaded,
total,
dlutil.GetSpeed(uploaded, p.start),
p.uploadAttempt,
), true)
}
log.FromContext(ctx).Debugf("Progress update: %s, %d/%d", info.FileName(), downloaded, total)
entityBuilder := entity.Builder{}
var entities []tg.MessageEntityClass
if err := styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressFileProcessingPrefix, nil)),
styling.Code(info.FileName()),
styling.Plain(i18n.T(i18nk.BotMsgProgressSavePathPrefix, nil)),
styling.Code(fmt.Sprintf("[%s]:%s", info.StorageName(), info.StoragePath())),
styling.Plain(i18n.T(i18nk.BotMsgProgressFileSizePrefix, nil)),
styling.Code(fmt.Sprintf("%.2f MB", float64(total)/(1024*1024))),
styling.Plain(i18n.T(i18nk.BotMsgProgressAvgSpeedPrefix, nil)),
styling.Bold(fmt.Sprintf("%.2f MB/s", dlutil.GetSpeed(downloaded, p.start)/(1024*1024))),
styling.Plain(i18n.T(i18nk.BotMsgProgressCurrentProgressPrefix, nil)),
styling.Bold(fmt.Sprintf("%.2f%%", float64(downloaded)/float64(total)*100)),
); err != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", err)
return
}
text, entities := entityBuilder.Complete()
req := &tg.MessagesEditMessageRequest{
ID: p.MessageID,
}
req.SetMessage(text)
req.SetEntities(entities)
req.SetReplyMarkup(&tg.ReplyInlineMarkup{
Rows: []tg.KeyboardButtonRow{
{
Buttons: []tg.KeyboardButtonClass{
tgutil.BuildCancelButton(info.TaskID()),
},
},
}},
)
ext := tgutil.ExtFromContext(ctx)
if ext != nil {
ext.EditMessage(p.ChatID, req)
return
}
func shouldUpdateUploadProgress(total, uploaded int64, lastPercent int, elapsed time.Duration) bool {
if total <= 0 || uploaded <= 0 {
return false
}
if uploaded >= total {
return lastPercent < 100 && elapsed >= uploadProgressMinInterval
}
percent := int((uploaded * 100) / total)
if percent < lastPercent {
return false
}
if elapsed < uploadProgressMinInterval {
return false
}
if percent == lastPercent {
return elapsed >= uploadProgressMaxInterval
}
return progressutil.ShouldUpdate(total, uploaded, lastPercent) || elapsed >= uploadProgressMaxInterval
}
func singleUploadPhase(attempt int) singleProgressPhase {
if attempt > 1 {
return singlePhaseRetrying
}
return singlePhaseUploading
}
func (p *Progress) OnDone(ctx context.Context, info TaskInfo, err error) {
p.updateMu.Lock()
defer p.updateMu.Unlock()
if err != nil {
log.FromContext(ctx).Errorf("Progress error for file [%s]: %v", info.FileName(), err)
} else {
log.FromContext(ctx).Debugf("Progress done for file [%s]", info.FileName())
}
p.editMessage(ctx, info.TaskID(), buildSingleDoneMessage(info, p.doneSize(info), err), false)
}
entityBuilder := entity.Builder{}
var stylingErr error
func (p *Progress) doneSize(info TaskInfo) int64 {
if p.hasActualSize {
return p.actualSize
if err != nil {
if errors.Is(err, context.Canceled) {
stylingErr = styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressTaskCanceled, nil)),
styling.Plain("\n"),
styling.Plain(i18n.T(i18nk.BotMsgProgressFileNamePrefix, nil)),
styling.Code(info.FileName()),
)
} else {
stylingErr = styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressDownloadFailedPrefix, nil)),
styling.Code(info.FileName()),
styling.Plain(i18n.T(i18nk.BotMsgProgressErrorPrefix, nil)),
styling.Bold(err.Error()),
)
}
} else {
stylingErr = styling.Perform(&entityBuilder,
styling.Plain(i18n.T(i18nk.BotMsgProgressDownloadDonePrefix, nil)),
styling.Code(info.FileName()),
styling.Plain(i18n.T(i18nk.BotMsgProgressSavePathPrefix, nil)),
styling.Code(fmt.Sprintf("[%s]:%s", info.StorageName(), info.StoragePath())),
)
}
return max(info.FileSize(), 0)
}
func (p *Progress) editMessage(ctx context.Context, taskID string, message renderedSingleMessage, cancellable bool) {
if message.Err != nil {
log.FromContext(ctx).Errorf("Failed to render file progress message: %v", message.Err)
if stylingErr != nil {
log.FromContext(ctx).Errorf("Failed to build entities: %s", stylingErr)
return
}
req := buildSingleEditMessageRequest(p.MessageID, taskID, message, cancellable)
if ext := tgutil.ExtFromContext(ctx); ext != nil {
if _, err := ext.EditMessage(p.ChatID, req); err != nil {
log.FromContext(ctx).Errorf("Failed to edit file progress message: %v", err)
}
}
}
func buildSingleEditMessageRequest(messageID int, taskID string, message renderedSingleMessage, cancellable bool) *tg.MessagesEditMessageRequest {
req := &tg.MessagesEditMessageRequest{ID: messageID}
req.SetMessage(message.Text)
if len(message.Entities) > 0 {
req.SetEntities(message.Entities)
text, entities := entityBuilder.Complete()
req := &tg.MessagesEditMessageRequest{
ID: p.MessageID,
}
if cancellable {
req.SetReplyMarkup(&tg.ReplyInlineMarkup{Rows: []tg.KeyboardButtonRow{{
Buttons: []tg.KeyboardButtonClass{tgutil.BuildCancelButton(taskID)},
}}})
}
return req
}
req.SetMessage(text)
req.SetEntities(entities)
func buildSingleProgressMessage(
info TaskInfo,
phase singleProgressPhase,
current int64,
total int64,
speed float64,
attempt int,
) renderedSingleMessage {
if current < 0 {
current = 0
ext := tgutil.ExtFromContext(ctx)
if ext != nil {
ext.EditMessage(p.ChatID, req)
}
if total > 0 && current > total {
current = total
}
percent := singleProgressPercent(current, total)
destination := fmt.Sprintf("[%s]:%s", info.StorageName(), info.StoragePath())
data := map[string]any{
"Name": info.FileName(),
"Bar": singleProgressBar(percent),
"Progress": percent,
"Speed": singleProgressSpeed(speed),
"Current": dlutil.FormatSize(current),
"Size": dlutil.FormatSize(total),
"Destination": destination,
"Attempt": max(attempt, 1),
}
var key i18nk.Key
switch phase {
case singlePhaseUploading:
key = i18nk.BotMsgProgressSingleUploading
case singlePhaseRetrying:
key = i18nk.BotMsgProgressSingleUploadRetrying
default:
key = i18nk.BotMsgProgressSingleDownloading
if total <= 0 {
key = i18nk.BotMsgProgressSingleDownloadingUnknown
}
}
markup := i18n.T(i18nk.BotMsgProgressSingleStatusHeader, nil) + "\n\n" + localizedProgressMarkup(key, data)
return completeSingleMessage(markup)
}
func buildSingleDoneMessage(info TaskInfo, size int64, err error) renderedSingleMessage {
data := map[string]any{
"Name": info.FileName(),
"Size": dlutil.FormatSize(max(size, 0)),
"Destination": fmt.Sprintf("[%s]:%s", info.StorageName(), info.StoragePath()),
}
var key i18nk.Key
switch {
case err == nil:
key = i18nk.BotMsgProgressSingleDone
case errors.Is(err, context.Canceled):
key = i18nk.BotMsgProgressSingleCanceled
default:
data["Reason"] = truncateSingleError(err.Error())
key = i18nk.BotMsgProgressSingleFailed
}
return completeSingleMessage(localizedProgressMarkup(key, data))
}
func localizedProgressMarkup(key i18nk.Key, data map[string]any) string {
return i18n.T(key, tgutil.EscapeHTMLTemplateData(data))
}
func completeSingleMessage(markup string) renderedSingleMessage {
text, entities, err := tgutil.RenderHTML(markup)
return renderedSingleMessage{Text: text, Entities: entities, Err: err}
}
func singleProgressPercent(current, total int64) int {
if total <= 0 {
return 0
}
return int(min(max(current, 0), total) * 100 / total)
}
func singleProgressBar(percent int) string {
percent = min(max(percent, 0), 100)
filled := percent * singleProgressBarWidth / 100
return strings.Repeat("🟩", filled) + strings.Repeat("⬜️", singleProgressBarWidth-filled)
}
func singleProgressSpeed(speed float64) string {
if speed <= 0 {
return "0 B/s"
}
return dlutil.FormatSize(int64(speed)) + "/s"
}
func truncateSingleError(value string) string {
runes := []rune(value)
if len(runes) <= maxSingleErrorRunes {
return value
}
return string(runes[:maxSingleErrorRunes])
}
type ProgressOption func(*Progress)

View File

@@ -1,178 +0,0 @@
package tfile
import (
"context"
"errors"
"strings"
"sync"
"testing"
"time"
"github.com/gotd/td/tg"
"github.com/krau/SaveAny-Bot/common/i18n"
)
type progressTestTaskInfo struct{}
func (progressTestTaskInfo) TaskID() string { return "task" }
func (progressTestTaskInfo) FileName() string { return "file.bin" }
func (progressTestTaskInfo) FileSize() int64 { return 100 << 20 }
func (progressTestTaskInfo) StoragePath() string { return "file.bin" }
func (progressTestTaskInfo) StorageName() string { return "test" }
func TestShouldUpdateUploadProgress(t *testing.T) {
tests := []struct {
name string
total int64
uploaded int64
lastPercent int
elapsed time.Duration
want bool
}{
{name: "invalid total", total: 0, uploaded: 1, want: false},
{name: "no uploaded bytes", total: 100, uploaded: 0, want: false},
{name: "percentage threshold", total: 100 << 20, uploaded: 10 << 20, elapsed: uploadProgressMinInterval, want: true},
{name: "percentage threshold rate limited", total: 100 << 20, uploaded: 10 << 20, elapsed: uploadProgressMinInterval - time.Millisecond, want: false},
{name: "maximum time threshold", total: 100 << 20, uploaded: 1 << 20, elapsed: uploadProgressMaxInterval, want: true},
{name: "below thresholds", total: 100 << 20, uploaded: 1 << 20, elapsed: uploadProgressMaxInterval - time.Millisecond, want: false},
{name: "completion", total: 100, uploaded: 100, lastPercent: 99, elapsed: uploadProgressMinInterval, want: true},
{name: "completion rate limited", total: 100, uploaded: 100, lastPercent: 99, elapsed: uploadProgressMinInterval - time.Millisecond, want: false},
{name: "completion already reported", total: 100, uploaded: 100, lastPercent: 100, elapsed: uploadProgressMinInterval, want: false},
{name: "out of order callback", total: 100, uploaded: 40, lastPercent: 60, elapsed: uploadProgressMaxInterval, want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := shouldUpdateUploadProgress(tt.total, tt.uploaded, tt.lastPercent, tt.elapsed)
if got != tt.want {
t.Fatalf("shouldUpdateUploadProgress() = %v, want %v", got, tt.want)
}
})
}
}
func TestUploadProgressConcurrentCallbacks(t *testing.T) {
progress := new(Progress)
ctx := context.Background()
info := progressTestTaskInfo{}
const total = int64(100 << 20)
progress.OnUploadStart(ctx, info, total)
progress.lastUpdateAt.Store(time.Now().Add(-uploadProgressMaxInterval).UnixNano())
var wg sync.WaitGroup
for uploaded := int64(1 << 20); uploaded <= total; uploaded += 1 << 20 {
uploaded := uploaded
wg.Go(func() {
progress.OnUploadProgress(ctx, info, uploaded, total)
})
}
wg.Wait()
percent := progress.lastUpdatePercent.Load()
if percent <= 0 || percent > 100 {
t.Fatalf("last upload percentage = %d, want a value in (0, 100]", percent)
}
if progress.uploadedBytes != total {
t.Fatalf("maximum uploaded bytes = %d, want %d", progress.uploadedBytes, total)
}
}
func TestSingleUploadRetryKeepsRichProgressLayout(t *testing.T) {
i18n.Init("zh-Hans")
t.Cleanup(func() { i18n.Init("zh-Hans") })
message := buildSingleProgressMessage(
progressTestTaskInfo{},
singleUploadPhase(2),
25<<20,
100<<20,
5<<20,
2,
)
if message.Err != nil {
t.Fatalf("buildSingleProgressMessage() failed: %v", message.Err)
}
for _, want := range []string{
"🔁 上传重试",
"🟩🟩⬜️⬜️⬜️⬜️⬜️⬜️⬜️⬜️ 25%",
"尝试次数2",
"速度5.00 MB/s",
"大小25.00 MB / 100.00 MB",
} {
if !strings.Contains(message.Text, want) {
t.Fatalf("retry progress does not contain %q:\n%s", want, message.Text)
}
}
}
func TestSingleProgressTemplateOwnsStylesAndEscapesValues(t *testing.T) {
i18n.Init("en")
t.Cleanup(func() { i18n.Init("zh-Hans") })
info := htmlProgressTestTaskInfo{}
message := buildSingleProgressMessage(info, singlePhaseDownloading, 50, 100, 25, 0)
if message.Err != nil {
t.Fatalf("buildSingleProgressMessage() failed: %v", message.Err)
}
for _, want := range []string{
`<b>A&B</b>.bin`,
`[store<&>]:dir/<i>x</i>&`,
"Speed: 25 B/s",
} {
if !strings.Contains(message.Text, want) {
t.Fatalf("progress text does not contain %q:\n%s", want, message.Text)
}
}
bold, code, blockquote, italic := singleEntityCounts(message.Entities)
if bold != 2 || code != 6 || blockquote != 1 || italic != 0 {
t.Fatalf("progress entity counts = bold:%d code:%d blockquote:%d italic:%d", bold, code, blockquote, italic)
}
failure := buildSingleDoneMessage(info, 100, errors.New(`<i>remote & failed</i>`))
if failure.Err != nil {
t.Fatalf("buildSingleDoneMessage() failed: %v", failure.Err)
}
if !strings.Contains(failure.Text, `<i>remote & failed</i>`) {
t.Fatalf("failure reason was not preserved literally:\n%s", failure.Text)
}
bold, code, blockquote, italic = singleEntityCounts(failure.Entities)
if bold != 1 || code != 2 || blockquote != 0 || italic != 0 {
t.Fatalf("failure entity counts = bold:%d code:%d blockquote:%d italic:%d", bold, code, blockquote, italic)
}
}
func TestSingleDoneSizeUsesActualUploadSize(t *testing.T) {
progress := new(Progress)
info := progressTestTaskInfo{}
progress.OnStart(context.Background(), info)
progress.OnUploadStart(context.Background(), info, 2048)
if got := progress.doneSize(info); got != 2048 {
t.Fatalf("done size = %d, want actual upload size 2048", got)
}
}
type htmlProgressTestTaskInfo struct{}
func (htmlProgressTestTaskInfo) TaskID() string { return "html-task" }
func (htmlProgressTestTaskInfo) FileName() string { return `<b>A&B</b>.bin` }
func (htmlProgressTestTaskInfo) FileSize() int64 { return 100 }
func (htmlProgressTestTaskInfo) StoragePath() string { return `dir/<i>x</i>&` }
func (htmlProgressTestTaskInfo) StorageName() string { return `store<&>` }
func singleEntityCounts(entities []tg.MessageEntityClass) (bold, code, blockquote, italic int) {
for _, messageEntity := range entities {
switch messageEntity.(type) {
case *tg.MessageEntityBold:
bold++
case *tg.MessageEntityCode:
code++
case *tg.MessageEntityBlockquote:
blockquote++
case *tg.MessageEntityItalic:
italic++
}
}
return
}

View File

@@ -23,8 +23,6 @@ type Task struct {
Progress ProgressTracker
stream bool // true if the file should be downloaded in stream mode
localPath string
overwrite bool // recovered: overwrite the storage target instead of uniquifying
caption string // recovered: source caption for the telegram backend
}
// Title implements core.Exectable.

32
core/tasks/tfile/util.go Normal file
View File

@@ -0,0 +1,32 @@
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
}

View File

@@ -1,75 +0,0 @@
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)
}
}

View File

@@ -1,102 +0,0 @@
package transfer
import (
"context"
"encoding/json"
"fmt"
"github.com/krau/SaveAny-Bot/core"
"github.com/krau/SaveAny-Bot/pkg/enums/ctxkey"
"github.com/krau/SaveAny-Bot/pkg/enums/tasktype"
"github.com/krau/SaveAny-Bot/pkg/storagetypes"
"github.com/krau/SaveAny-Bot/storage"
)
func ctxOverwrite(ctx context.Context) bool {
overwrite, _ := ctx.Value(ctxkey.OverwriteExisting).(bool)
return overwrite
}
type elementPayload struct {
ID string `json:"id"`
SourceStorage string `json:"source_storage"`
SourcePath string `json:"source_path"`
FileInfo storagetypes.FileInfo `json:"file_info"`
TargetStorage string `json:"target_storage"`
TargetPath string `json:"target_path"`
}
type taskPayload struct {
ID string `json:"id"`
Elements []elementPayload `json:"elements"`
ChatID int64 `json:"chat_id"`
MessageID int `json:"message_id"`
IgnoreErrors bool `json:"ignore_errors"`
Overwrite bool `json:"overwrite"`
}
type taskCodec struct{}
func init() {
core.RegisterTaskCodec(tasktype.TaskTypeTransfer, taskCodec{})
}
func (taskCodec) Marshal(task core.Executable) ([]byte, error) {
t, ok := task.(*Task)
if !ok {
return nil, fmt.Errorf("unexpected task type %T", task)
}
p := taskPayload{
ID: t.ID,
IgnoreErrors: t.IgnoreErrors,
Overwrite: ctxOverwrite(t.ctx),
}
for _, elem := range t.elems {
p.Elements = append(p.Elements, elementPayload{
ID: elem.ID,
SourceStorage: elem.SourceStorage.Name(),
SourcePath: elem.SourcePath,
FileInfo: elem.FileInfo,
TargetStorage: elem.TargetStorage.Name(),
TargetPath: elem.TargetPath,
})
}
if progress, ok := t.Progress.(*Progress); ok {
p.ChatID = progress.ChatID
p.MessageID = progress.MessageID
}
return json.Marshal(p)
}
func (taskCodec) Unmarshal(data []byte) (core.Executable, error) {
var p taskPayload
if err := json.Unmarshal(data, &p); err != nil {
return nil, fmt.Errorf("invalid task payload: %w", err)
}
elems := make([]TaskElement, 0, len(p.Elements))
for _, ep := range p.Elements {
source, err := storage.GetStorageByName(context.Background(), ep.SourceStorage)
if err != nil {
return nil, fmt.Errorf("source storage %q: %w", ep.SourceStorage, err)
}
target, err := storage.GetStorageByName(context.Background(), ep.TargetStorage)
if err != nil {
return nil, fmt.Errorf("target storage %q: %w", ep.TargetStorage, err)
}
elems = append(elems, TaskElement{
ID: ep.ID,
SourceStorage: source,
SourcePath: ep.SourcePath,
FileInfo: ep.FileInfo,
TargetStorage: target,
TargetPath: ep.TargetPath,
})
}
var progress ProgressTracker
if p.ChatID != 0 {
progress = NewProgressTracker(p.MessageID, p.ChatID)
}
task := NewTransferTask(p.ID, context.Background(), elems, progress, p.IgnoreErrors)
task.overwrite = p.Overwrite
return task, nil
}

View File

@@ -2,7 +2,6 @@ package transfer
import (
"context"
"errors"
"fmt"
"io"
"os"
@@ -21,12 +20,7 @@ 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")
if t.overwrite {
ctx = storage.WithOverwrite(ctx)
}
if t.Progress != nil {
t.Progress.OnStart(ctx, t)
}
t.Progress.OnStart(ctx, t)
workers := config.C().Workers
eg, gctx := errgroup.WithContext(ctx)
@@ -34,11 +28,14 @@ func (t *Task) Execute(ctx context.Context) error {
for _, elem := range t.elems {
eg.Go(func() error {
t.processingMu.Lock()
t.processingMu.RLock()
if t.processing[elem.ID] != nil {
t.processingMu.Unlock()
t.processingMu.RUnlock()
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()
@@ -49,7 +46,7 @@ func (t *Task) Execute(ctx context.Context) error {
}()
err := t.processElement(gctx, elem)
if err != nil && (!t.IgnoreErrors || errors.Is(err, context.Canceled)) {
if err != nil && !t.IgnoreErrors {
return err
}
if err != nil {
@@ -69,9 +66,7 @@ func (t *Task) Execute(ctx context.Context) error {
logger.Info("Transfer task completed successfully")
}
if t.Progress != nil {
t.Progress.OnDone(ctx, t, err)
}
t.Progress.OnDone(ctx, t, err)
return err
}
@@ -121,9 +116,7 @@ func (t *Task) processElement(ctx context.Context, elem TaskElement) error {
}
t.uploaded.Add(size)
if t.Progress != nil {
t.Progress.OnProgress(ctx, t)
}
t.Progress.OnProgress(ctx, t)
taskevent.Emit(ctx, taskevent.Event{
TaskID: t.ID,
Phase: taskevent.PhaseProgress,

View File

@@ -14,7 +14,6 @@ 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"
)
@@ -84,7 +83,7 @@ func (p *Progress) OnStart(ctx context.Context, info TaskInfo) {
}
func (p *Progress) OnProgress(ctx context.Context, info TaskInfo) {
if !progressutil.ShouldUpdate(info.TotalSize(), info.Uploaded(), int(p.lastUpdatePercent.Load())) {
if !shouldUpdateProgress(info.TotalSize(), info.Uploaded(), int(p.lastUpdatePercent.Load())) {
return
}
percent := int((info.Uploaded() * 100) / info.TotalSize())
@@ -222,6 +221,14 @@ 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

View File

@@ -35,7 +35,6 @@ type Task struct {
processing map[string]TaskElementInfo
processingMu sync.RWMutex
failed map[string]error
overwrite bool // recovered: overwrite storage targets instead of uniquifying
}
// Title implements core.Executable.

View File

@@ -35,7 +35,7 @@ func Init(ctx context.Context) {
logger.Fatal("Failed to open database: ", err)
}
logger.Debug("Database connected")
if err := db.AutoMigrate(&User{}, &Dir{}, &Rule{}, &WatchChat{}, &Task{}); err != nil {
if err := db.AutoMigrate(&User{}, &Dir{}, &Rule{}, &WatchChat{}); err != nil {
logger.Fatal("Database migration failed; if upgrading from an old version, try deleting the database file and retrying", "error", err)
}
if err := syncUsers(ctx); err != nil {

View File

@@ -1,137 +0,0 @@
package database
import (
"context"
"errors"
"time"
)
var errNotInitialized = errors.New("database not initialized")
type TaskStatus string
const (
TaskStatusQueued TaskStatus = "queued"
TaskStatusRunning TaskStatus = "running"
TaskStatusFailed TaskStatus = "failed"
TaskStatusCancelled TaskStatus = "cancelled"
)
// Task is the persisted record of a queued or running task, used to recover
// unfinished work after a process restart. Completed tasks are deleted on
// finish, so the table only ever holds queued/running rows.
type Task struct {
ID string `gorm:"primaryKey;size:64"`
Type string `gorm:"size:32;index"`
Payload []byte
Status string `gorm:"size:16;index"`
Error string
CreatedAt time.Time
UpdatedAt time.Time
}
func CreateTask(ctx context.Context, task *Task) error {
if db == nil {
return errNotInitialized
}
return db.WithContext(ctx).Create(task).Error
}
// UpsertTask inserts the task or replaces the existing row with the same ID.
func UpsertTask(ctx context.Context, task *Task) error {
if db == nil {
return errNotInitialized
}
return db.WithContext(ctx).Save(task).Error
}
func UpdateTaskStatus(ctx context.Context, id string, status TaskStatus, errMsg string) error {
if db == nil {
return errNotInitialized
}
return db.WithContext(ctx).Model(&Task{}).
Where("id = ?", id).
Updates(map[string]any{
"status": status,
"error": errMsg,
"updated_at": time.Now(),
}).Error
}
func GetTask(ctx context.Context, id string) (*Task, error) {
if db == nil {
return nil, errNotInitialized
}
var task Task
if err := db.WithContext(ctx).First(&task, "id = ?", id).Error; err != nil {
return nil, err
}
return &task, nil
}
// UpdateTaskPayload replaces the payload of an existing task row.
func UpdateTaskPayload(ctx context.Context, id string, payload []byte) error {
if db == nil {
return errNotInitialized
}
return db.WithContext(ctx).Model(&Task{}).
Where("id = ?", id).
Updates(map[string]any{
"payload": payload,
"updated_at": time.Now(),
}).Error
}
// RestoreTaskCreatedAt restores the original creation time after a
// re-enqueue overwrote it.
func RestoreTaskCreatedAt(ctx context.Context, id string, createdAt time.Time) error {
if db == nil {
return errNotInitialized
}
return db.WithContext(ctx).Model(&Task{}).
Where("id = ?", id).
Update("created_at", createdAt).Error
}
// DeleteStaleFailedTasks removes failed rows older than the given age.
func DeleteStaleFailedTasks(ctx context.Context, maxAge time.Duration) error {
if db == nil {
return errNotInitialized
}
return db.WithContext(ctx).
Where("status = ? AND updated_at < ?", string(TaskStatusFailed), time.Now().Add(-maxAge)).
Delete(&Task{}).Error
}
func DeleteTask(ctx context.Context, id string) error {
if db == nil {
return errNotInitialized
}
return db.WithContext(ctx).Delete(&Task{}, "id = ?", id).Error
}
// GetUnfinishedTasks returns all tasks that were not finished when the
// process stopped, i.e. tasks that must be re-enqueued on startup.
func GetUnfinishedTasks(ctx context.Context) ([]Task, error) {
if db == nil {
return nil, errNotInitialized
}
var tasks []Task
err := db.WithContext(ctx).
Where("status IN ?", []string{string(TaskStatusQueued), string(TaskStatusRunning)}).
Order("created_at").
Find(&tasks).Error
return tasks, err
}
func CountUnfinishedTasks(ctx context.Context) (int64, error) {
if db == nil {
return 0, errNotInitialized
}
var count int64
err := db.WithContext(ctx).
Model(&Task{}).
Where("status IN ?", []string{string(TaskStatusQueued), string(TaskStatusRunning)}).
Count(&count).Error
return count, err
}

View File

@@ -1,110 +0,0 @@
package database
import (
"context"
"path/filepath"
"testing"
"github.com/ncruces/go-sqlite3/gormlite"
"gorm.io/gorm"
)
func newTestDB(t *testing.T) {
t.Helper()
d, err := gorm.Open(gormlite.Open(filepath.Join(t.TempDir(), "test.db")), &gorm.Config{})
if err != nil {
t.Fatalf("open test db: %v", err)
}
if err := d.AutoMigrate(&Task{}); err != nil {
t.Fatalf("migrate: %v", err)
}
old := db
db = d
t.Cleanup(func() { db = old })
}
func TestTaskCRUD(t *testing.T) {
newTestDB(t)
ctx := context.Background()
task := &Task{
ID: "task-1",
Type: "tfile",
Payload: []byte(`{"file":"x"}`),
Status: string(TaskStatusQueued),
}
if err := CreateTask(ctx, task); err != nil {
t.Fatalf("create: %v", err)
}
unfinished, err := GetUnfinishedTasks(ctx)
if err != nil {
t.Fatalf("get unfinished: %v", err)
}
if len(unfinished) != 1 || unfinished[0].ID != "task-1" {
t.Fatalf("got %+v, want 1 task task-1", unfinished)
}
if err := UpdateTaskStatus(ctx, "task-1", TaskStatusRunning, ""); err != nil {
t.Fatalf("update: %v", err)
}
unfinished, err = GetUnfinishedTasks(ctx)
if err != nil {
t.Fatalf("get unfinished after update: %v", err)
}
if len(unfinished) != 1 || unfinished[0].Status != string(TaskStatusRunning) {
t.Fatalf("running status not persisted: %+v", unfinished)
}
if err := DeleteTask(ctx, "task-1"); err != nil {
t.Fatalf("delete: %v", err)
}
count, err := CountUnfinishedTasks(ctx)
if err != nil {
t.Fatalf("count: %v", err)
}
if count != 0 {
t.Fatalf("count = %d, want 0", count)
}
}
func TestTaskUpsert(t *testing.T) {
newTestDB(t)
ctx := context.Background()
task := &Task{ID: "task-2", Type: "tfile", Status: string(TaskStatusQueued)}
if err := UpsertTask(ctx, task); err != nil {
t.Fatalf("upsert create: %v", err)
}
task.Status = string(TaskStatusRunning)
task.Payload = []byte("new")
if err := UpsertTask(ctx, task); err != nil {
t.Fatalf("upsert update: %v", err)
}
unfinished, err := GetUnfinishedTasks(ctx)
if err != nil {
t.Fatalf("get unfinished: %v", err)
}
if len(unfinished) != 1 || unfinished[0].Status != string(TaskStatusRunning) || string(unfinished[0].Payload) != "new" {
t.Fatalf("upsert did not replace: %+v", unfinished)
}
}
func TestGetUnfinishedTasksExcludesFinished(t *testing.T) {
newTestDB(t)
ctx := context.Background()
if err := CreateTask(ctx, &Task{ID: "done", Type: "tfile", Status: string(TaskStatusFailed)}); err != nil {
t.Fatal(err)
}
if err := CreateTask(ctx, &Task{ID: "pending", Type: "tfile", Status: string(TaskStatusQueued)}); err != nil {
t.Fatal(err)
}
unfinished, err := GetUnfinishedTasks(ctx)
if err != nil {
t.Fatal(err)
}
if len(unfinished) != 1 || unfinished[0].ID != "pending" {
t.Fatalf("got %+v, want only pending", unfinished)
}
}

View File

@@ -113,51 +113,6 @@ secret = "your-rpc-secret"
remove_after_transfer = true
```
### yt-dlp Configuration
Configures the behavior of the `/ytdlp` command and the `ytdlp` HTTP-API task type when no custom flags are passed.
- `max_height`: Default maximum video resolution by height in pixels (e.g. `1080`, `720`). `0` means no limit (best available). Ignored when `format` is set.
- `format`: A raw yt-dlp format selector (`-f`). When set, it takes precedence over `max_height` and gives you full control, e.g. `bv*[height<=720]+ba/b`.
- `recode`: The target video container yt-dlp recodes into after download (e.g. `mp4`). Leave empty to disable recoding.
{{< hint info >}}
These defaults only apply when using the `/ytdlp` command without passing any custom flags. Passing custom flags on the command (or `flags` in the API) overrides them.
{{< /hint >}}
```toml
[ytdlp]
max_height = 1080
format = "" # e.g. "bv*[height<=720]+ba/b"
recode = "mp4" # empty disables recoding
```
### HTTP API Configuration
When enabled, SaveAny-Bot exposes an HTTP API for creating/querying/canceling tasks programmatically. See [HTTP API](../../usage/api) for the full endpoint reference.
- `enable`: Whether to enable the HTTP API server, default is `false`.
- `host`: Bind address, default `0.0.0.0`.
- `port`: Listen port, default `8080`.
- `token`: Authentication token. **Strongly recommended** — if empty, the API is exposed without any authentication.
```toml
[api]
enable = false
host = "0.0.0.0"
port = 8080
token = "your-token"
```
### Log Configuration
- `level`: Log level. One of `debug`, `info`, `warn`, `error`, `fatal`. Default is `info`.
```toml
[log]
level = "info"
```
### Storage Endpoints List
The storage endpoints list is used to define the storage locations supported by the Bot. Each storage endpoint needs to specify a name, type, and related configuration, using the double bracket syntax `[[storages]]`.

View File

@@ -79,8 +79,7 @@ Stream mode is not supported.
chat_id = "123456789" # Telegram chat ID, the bot will send files to this chat
force_file = false # Force sending as file, default is false
skip_large = false # Skip large files, default is false. If enabled, files exceeding Telegram's limit will not be uploaded.
split_large_video = false # Losslessly split oversized videos into one album of up to 10 independently playable parts. Falls back to ZIP parts on failure or when more than 10 parts are required.
split_size_mb = 0 # Split size in MB. 0 uses the uploader account limit: 2000 MB for bots/regular users and 4000 MB for Premium users. Oversized non-video files use ZIP parts. Ignored when skip_large is true.
spilt_size_mb = 2000 # Split size in MB, default is 2000 MB (2 GB). Files larger than this will be split into multiple parts (zip format). Ignored when skip_large is true.
```
## Rclone
@@ -137,4 +136,4 @@ remote = "myremote"
base_path = "/backup"
config_path = "/path/to/rclone.conf"
flags = ["--progress"]
```
```

View File

@@ -1,90 +0,0 @@
---
title: "CLI Subcommands"
weight: 21
---
# CLI Subcommands
Besides running the Telegram bot with `./saveany-bot` (no subcommand), the binary exposes two helper subcommands for moving local files into a storage backend: `upload` (one-shot) and `watch` (continuous).
These subcommands load the same `config.toml` as the bot, initialize the database and caches, then perform their task. They do **not** start the Telegram bot itself, although storages of type `telegram` will spin up the bot client just for the upload.
## `upload` — Upload a Single File
```
saveany-bot upload -f <file> -s <storage> [-d <dir>] [--no-progress]
```
Flags:
| Flag | Required | Description |
|---|---|---|
| `-f, --file` | Yes | Path to the local file to upload |
| `-s, --storage` | Yes | Target storage name (must exist in `config.toml`) |
| `-d, --dir` | No | Destination directory within the storage. Defaults to the storage's `base_path` |
| `--no-progress` | No | Disable the terminal progress bar |
Examples:
```bash
# Upload a file to the default dir of storage "MyAlist"
./saveany-bot upload -f ./movie.mp4 -s MyAlist
# Upload into a specific subdirectory
./saveany-bot upload -f ./movie.mp4 -s MyAlist -d movies/2026
# Upload via Telegram storage without a progress bar
./saveany-bot upload -f ./photo.jpg -s MyChannel --no-progress
```
## `watch` — Watch a Directory and Auto-Upload
The `watch` subcommand continuously monitors a local directory and uploads created or modified files to a storage backend, preserving the relative directory structure from the watch root.
```
saveany-bot watch -p <path> -s <storage> [-d <dir>] [options]
```
Flags:
| Flag | Default | Description |
|---|---|---|
| `-p, --path` | *(required)* | Local directory to watch |
| `-s, --storage` | *(required)* | Target storage name |
| `-d, --dir` | storage's `base_path` | Destination directory within the storage |
| `-r, --recursive` | `false` | Watch subdirectories recursively |
| `--overwrite` | `false` | Overwrite existing files on the storage instead of skipping them |
| `--initial-scan` | `false` | Upload files already present in the directory on startup |
| `--debounce` | `2s` | How long to wait after the last write before uploading a file |
| `--upload-workers` | `config.workers` | Number of concurrent uploads |
| `--retry-delay` | `3s` | Delay between upload retries |
{{< hint info >}}
Write-completion detection: the watcher debounces per file and only uploads once the file size stays unchanged across the debounce window, so partial/write-in-progress files are not uploaded.
<br />
If a file changes while being uploaded, it is re-uploaded once after the current upload finishes (instead of being queued multiple times).
{{< /hint >}}
Examples:
```bash
# Watch ./inbox and upload new files to "MyAlist" recursively
./saveany-bot watch -p ./inbox -s MyAlist -r
# Watch with a custom destination dir and overwrite
./saveany-bot watch -p ./inbox -s MyAlist -d backup --overwrite
# On startup, also upload everything already in ./inbox
./saveany-bot watch -p ./inbox -s MyAlist --initial-scan
```
### Behavior notes
- Relative directory structure is preserved under the destination directory. A file written to `./inbox/sub/file.txt` with `--path ./inbox` is uploaded to `<dest_dir>/sub/file.txt`.
- `watch` runs until interrupted (e.g. `Ctrl-C` / `SIGINT`); in-flight uploads are drained before exit.
- Retries follow the global `retry` value from `config.toml`, with `--retry-delay` between attempts.
- Telegram-type storages will start the bot client automatically to perform uploads.
{{< hint warning >}}
`watch` is unrelated to the in-bot `/watch` command (which watches Telegram chats). This subcommand watches a **local filesystem directory** and uploads to a storage backend, independent of Telegram.
{{< /hint >}}

View File

@@ -1,75 +0,0 @@
---
title: "File Naming & Conflict Strategies"
weight: 11
---
# File Naming & Conflict Strategies
SaveAny-Bot lets you customize how saved files are named and how collisions with existing files are resolved, directly in Telegram via the `/config` and `/fnametmpl` commands.
## `/config` — User Configuration
The `/config` command opens an inline menu where you can change two per-user settings:
- **Filename strategy** — how the saved file is named
- **Duplicate file strategy** — what happens when a file with the same name already exists in the target storage
Settings are stored per user and apply to all of that user's subsequent save/transfer tasks.
### Filename strategy
| Option | Behavior |
|---|---|
| `Default` | Use the original media filename, or a generated name when no original filename is available |
| `Gen From Msg First` | Generate the filename from the message content (e.g. caption, text) and prefer that over the original filename |
| `Template` | Render the filename from a custom template you define with `/fnametmpl` |
### Duplicate file strategy
| Option | Behavior |
|---|---|
| `Always rename` (default) | Keep the existing file and save the new one with an alternate name |
| `Ask every time` | Prompt you with inline buttons each time a collision occurs |
| `Always overwrite` | Replace the existing file with the new one |
| `Always skip` | Do nothing for conflicting files |
{{< hint info >}}
The conflict strategy only kicks in for storage backends that can detect the existence of a file. Backends that do not support existence checks will fall back to overwriting.
{{< /hint >}}
## `/fnametmpl` — Custom Filename Template
When the filename strategy is set to `Template`, SaveAny-Bot renders each saved file's name using the template configured via `/fnametmpl`.
```
/fnametmpl [template]
```
- Running `/fnametmpl` without arguments shows your current template and the help text.
- Running it with a template string sets that template as your filename template.
The template uses Go [`text/template`](https://pkg.go.dev/text/template) syntax. The available variables are:
| Variable | Description |
|---|---|
| `{{.msgid}}` | Telegram message ID |
| `{{.msgtags}}` | Hashtags found in the message, joined with `_` |
| `{{.msggen}}` | Filename generated from the message |
| `{{.msgdate}}` | Message date, formatted `YYYY-MM-DD_HH-MM-SS` |
| `{{.msgraw}}` | Raw, unprocessed message text |
| `{{.origname}}` | The media's original filename (if any) |
| `{{.chatid}}` | Chat ID of the message |
Examples:
```
# Fixed prefix + message id + date
/fnametmpl Image_{{.msgid}}_{{.msgdate}}.jpg
# Use original name if available, otherwise a generated name
/fnametmpl {{.origname}}
```
{{< hint warning >}}
The template only takes effect when the filename strategy is set to `Template`. If template parsing fails, SaveAny-Bot falls back to the default filename naming logic.
{{< /hint >}}

Some files were not shown because too many files have changed in this diff Show More