mirror of
https://github.com/krau/SaveAny-Bot.git
synced 2026-08-27 19:30:07 +08:00
Compare commits
1 Commits
feat/resum
...
refactor/p
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
df3c568bb8 |
4
.github/workflows/build-release.yml
vendored
4
.github/workflows/build-release.yml
vendored
@@ -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
1
.gitignore
vendored
@@ -1,6 +1,5 @@
|
||||
config.toml
|
||||
logs/
|
||||
/cache/
|
||||
tmp/
|
||||
data/
|
||||
downloads/
|
||||
|
||||
376
AGENTS.md
376
AGENTS.md
@@ -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
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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()),
|
||||
|
||||
@@ -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) {}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -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/")
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}()
|
||||
}
|
||||
|
||||
|
||||
@@ -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)})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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)))
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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)),
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
27
cmd/run.go
27
cmd/run.go
@@ -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)
|
||||
|
||||
33
common/cache/ristretto.go
vendored
33
common/cache/ristretto.go
vendored
@@ -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 {
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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"
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"], "<b>A&B</b>"; 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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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"`
|
||||
// }
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
49
core/core.go
49
core/core.go
@@ -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 {
|
||||
|
||||
145
core/persist.go
145
core/persist.go
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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()), " ")
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
32
core/tasks/batchtfile/utils.go
Normal file
32
core/tasks/batchtfile/utils.go
Normal 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
|
||||
}
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
|
||||
13
core/tasks/telegraph/utils.go
Normal file
13
core/tasks/telegraph/utils.go
Normal 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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
32
core/tasks/tfile/util.go
Normal 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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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 {
|
||||
|
||||
137
database/task.go
137
database/task.go
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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]]`.
|
||||
|
||||
@@ -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"]
|
||||
```
|
||||
```
|
||||
@@ -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 >}}
|
||||
@@ -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
Reference in New Issue
Block a user