diff --git a/CHANGELOG.md b/CHANGELOG.md index cbe7c5b..1e29149 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,15 +12,19 @@ - Added `PollAnswer.VoterUser` and `PollAnswer.VoterChatInfo` presence helpers without changing the v1 value-field layout. - Added typed runtime errors for invalid scene actions, nil handlers, recovered handler panics, oversized Telegram responses, and recovered worker-pool task panics. - Added `CommandScopeUpdateError` and command-generation validation errors so callers can inspect invalid commands, duplicate names, invalid descriptions, and partially updated Telegram scopes. +- Added rich-message editing helpers `MessageContext.EditCallbackRich`, `MessageContext.UpsertKeyboardRich`, `AnswerMessage.EditRich`, and `AnswerMessage.EditRichKeyboard`, plus `AnswerMessage.RichHTML` for rendered rich content. +- Added `NewBotWithAPI` for constructing a bot with a preconfigured, injectable Telegram API client. ### Changed -- Runtime observer callbacks now execute asynchronously in enqueue order on one bounded dispatcher instead of blocking update handlers and scene locks. Queue overflow drops events with sampled warnings, observer panics remain isolated, and runtime shutdown drains queued events. +- Runtime observer callbacks now execute asynchronously in enqueue order on one bounded dispatcher instead of blocking update handlers and scene locks. Queue overflow drops events with sampled warnings, observer panics remain isolated, and shutdown cancels callback contexts, drains queued events, and reports callbacks that ignore cancellation. +- Scene routing now holds a lock only for the active session key; updates without an active scene and unrelated user sessions in the same chat remain concurrent. - Multipart uploads now encode request bodies as streams instead of building a second complete in-memory copy for every attempt. - Logger replacement and app-data logger writer registration now follow the bot configuration freeze, reject nil values safely, and install bot-token redaction before publishing replacement loggers. -- README and README_RU now document context runners, bounded retries and downloads, streaming multipart uploads, keyboard validation, and asynchronous observer delivery. +- README and README_RU now document context runners, injectable API clients, bounded retries and downloads, streaming multipart uploads, keyboard validation, and asynchronous observer delivery. - Version metadata now reports `v1.2.0`. ### Fixed +- Fixed HTTP transport errors exposing the bot token through Bot API request and file-download URLs while preserving error-chain inspection. - Fixed per-chat rate-limiter lifecycle races with concurrent cleanup, cooldown waits ignoring a later extension, shorter 429 locks replacing longer locks, and busy chats consuming global capacity before obtaining chat capacity. - Fixed scene `AsUserError` values not reaching users and producing duplicate observer errors; scene step, command, payload, and message-fallback errors now follow one consistent error path. - Fixed nil and panicking command, payload, update, message-fallback, scene, and middleware callbacks breaking handler/observer lifecycles. Panics are recovered per invocation and matching finish/error events are emitted. @@ -35,16 +39,26 @@ - Fixed worker-pool task panics terminating a worker and leaving the caller without a result. - Fixed zero or negative Telegram `retry_after` values causing a tight retry loop. - Fixed Telegram JSON wire names for `provider_payment_charge_id`, `owned_gift_id`, `others_can_add_tasks`, `others_can_mark_tasks_as_done`, `available_reactions`, and `quote_parse_mode`. +- Fixed `MessageContext` text, photo, caption, rich-message, and chat-action helpers omitting business connection identifiers; message-backed helpers now also reject missing chats instead of panicking, while inline edits remain supported without a chat-backed message. +- Fixed chat-admin, chat-creator, and bot-admin policies ignoring cancellation and deadlines from the active `MessageContext`. +- Fixed command parsing treating tabs and newlines as part of command names and matching addressed bot usernames case-sensitively. +- Fixed nested `RichTextArray` values bypassing rich-message depth validation and renderer limits. +- Fixed polling panics being reported only through telemetry while `RunWithContext` returned success. +- Fixed asynchronous middleware outliving the bot runtime and normal middleware denials being reported as internal handler errors. +- Fixed caller-owned or aliased bot loggers being closed by the framework, bot-owned aliases being closed more than once, and replaced internal loggers leaking resources. +- Fixed configuration codecs and JSON option encoding panicking on nil inputs. ### Documentation - Added field-level Godoc for the complete exported API surface, including official Telegram Bot API optionality, ranges, formats, and field semantics. The repository-wide AST audit now reports no undocumented exported declarations or struct fields. +- Corrected draft-provider Godoc to avoid promising cryptographic unpredictability from `math/rand/v2` IDs. - Marked v1 compatibility surfaces that may change in v2: `PollAnswer` voter pointers, context-required runners, error-returning keyboard builders, strict rich-message parsing, lossless unknown rich types, bounded downloads, compatibility aliases, and corrected public field names. +- Clarified that drop-overflow mode rejects outgoing API requests, one-shot asynchronous runners are awaited during shutdown, and pending uploads are drained by the shared API client. ### CI - Pinned the lint workflow to Go `1.26.6`, `golangci/golangci-lint-action@v9.0.0`, and golangci-lint `v2.12.2` for reproducible analysis. ### Tests -- Added regression coverage for limiter concurrency and cooldown behavior, all scene handler error paths, handler and worker panics, asynchronous observer delivery and drain, runner cancellation, draft cleanup, command validation and partial scopes, keyboard byte boundaries and rows, MarkdownV2 vectors, numeric binding boundaries, rich JSON limits and strict roots, bounded file downloads, retry caps, streaming multipart replay, and corrected Telegram wire keys. +- Added regression coverage for limiter concurrency and cooldown behavior, scene lock scope, all scene handler error paths, polling, handler, and worker panics, asynchronous observer delivery, cancellation, and bounded shutdown, asynchronous middleware lifecycle, logger ownership, runner and policy cancellation, nil configuration inputs, draft cleanup, command parsing, command validation and partial scopes, keyboard byte boundaries and rows, MarkdownV2 vectors, numeric binding boundaries, rich JSON and renderer depth limits, strict rich roots, rich-message editing, business-message helpers and photo replacement, bounded file downloads, retry caps, streaming multipart replay, and corrected Telegram wire keys. ## v1.1.0 diff --git a/README.md b/README.md index c6dd3d6..f16552c 100644 --- a/README.md +++ b/README.md @@ -120,6 +120,8 @@ func main() { 9. `RunWebhookWithContext(...)`: Starts the bot-owned webhook runtime when Telegram should deliver updates over HTTP instead of long polling. 10. A `Bot` instance is single-use. After `Run()`, `RunWithContext()`, or `RunWebhookWithContext()` returns, create a new bot instance for the next session. +For tests or custom transports, use `NewBotWithAPI[T](opts, api)` with a preconfigured `*tgapi.API`. The bot takes ownership of that client and closes it from `Bot.Close`; API transport, retry, and rate-limit fields in `BotOpts` do not override the supplied client. + ## File-Based Config `BotOpts` can also be loaded from or saved to config files through the file codec API. diff --git a/README_RU.md b/README_RU.md index bfc8dc9..e102e47 100644 --- a/README_RU.md +++ b/README_RU.md @@ -121,6 +121,8 @@ func main() { 9. `RunWebhookWithContext(...)`: Запускает bot-owned webhook runtime, когда Telegram должен доставлять update по HTTP вместо long polling. 10. Экземпляр `Bot` одноразовый. После завершения `Run()`, `RunWithContext()` или `RunWebhookWithContext()` для следующего запуска создавайте новый бот. +Для тестов или собственного transport используй `NewBotWithAPI[T](opts, api)` с заранее настроенным `*tgapi.API`. Бот становится владельцем этого клиента и закрывает его в `Bot.Close`; настройки transport, retry и rate limit из `BotOpts` не переопределяют переданный клиент. + ## Конфиг из файла `BotOpts` можно не только собирать вручную или из environment, но и загружать и сохранять через file codec API. diff --git a/bot.go b/bot.go index c980274..4e40828 100644 --- a/bot.go +++ b/bot.go @@ -97,13 +97,17 @@ type Bot[T AppData] struct { maxWorkers int pollTimeout int // Long-polling timeout in seconds for getUpdates - logFormat utils.LogFormat - logFormatter *sneklog.Formatter - logger *sneklog.Logger // Main bot logger (JSON stdout + optional file) - requestLogger *sneklog.Logger // Optional request-level API logging - useReqLogger bool - webhookLogger *sneklog.Logger // Webhook logger. Available only after Bot.RunWebhookWithContext. - extraLoggers extypes.Slice[*sneklog.Logger] // API, Uploader, and custom loggers + logFormat utils.LogFormat + logFormatter *sneklog.Formatter + logger *sneklog.Logger // Main bot logger (JSON stdout + optional file) + requestLogger *sneklog.Logger // Optional request-level API logging + useReqLogger bool + webhookLogger *sneklog.Logger // Webhook logger. Available only after Bot.RunWebhookWithContext. + loggerOwned bool + requestLoggerOwned bool + webhookLoggerOwned bool + detachedOwnedLoggers []*sneklog.Logger + extraLoggers extypes.Slice[*sneklog.Logger] // API, Uploader, and custom loggers plugins []Plugin[T] // Command/event handlers middlewares []Middleware[T] // Pre-processing filters (sorted by order) @@ -131,6 +135,7 @@ type Bot[T AppData] struct { updateQueue chan *tgapi.Update // Internal queue for processing updates runnerOnceWG sync.WaitGroup // Tracks one-time async runners runnerBgWG sync.WaitGroup // Tracks background async runners + middlewareWG sync.WaitGroup // Tracks asynchronous middleware callbacks runStateMu sync.Mutex running bool ran bool @@ -157,6 +162,20 @@ func (bot *Bot[T]) configMutable(method string) bool { // - Sets up DraftProvider with random IDs // - Adds API and Uploader loggers to extraLoggers func NewBot[T any](opts *BotOpts) (*Bot[T], error) { + return newBot[T](opts, nil) +} + +// NewBotWithAPI creates a Bot using a preconfigured API client. +// The Bot takes ownership of api and closes it from Bot.Close. API transport, +// retry, and rate-limit fields in opts do not reconfigure the supplied client. +func NewBotWithAPI[T any](opts *BotOpts, api *tgapi.API) (*Bot[T], error) { + if api == nil { + return nil, ErrAPIIsNil + } + return newBot[T](opts, api) +} + +func newBot[T any](opts *BotOpts, api *tgapi.API) (*Bot[T], error) { if opts == nil { return nil, ErrOptsIsNil } @@ -182,15 +201,17 @@ func NewBot[T any](opts *BotOpts) (*Bot[T], error) { // HTTP client timeout must exceed pollTimeout to avoid spurious deadline // errors that the polling loop would misinterpret as context cancellation. httpTimeout := time.Duration(pollTimeout)*time.Second + 60*time.Second - apiOpts := tgapi.NewAPIOpts(opts.Token). - SetAPIURL(opts.APIURL). - UseTestServer(opts.UseTestServer). - SetLimiter(limiter). - SetDropRateLimitOverflow(opts.DropRateLimitOverflow). - SetLogFormat(opts.LogFormat). - SetLogFormatter(opts.LogFormatter). - SetHTTPClient(&http.Client{Timeout: httpTimeout}) - api := tgapi.NewAPI(apiOpts) + if api == nil { + apiOpts := tgapi.NewAPIOpts(opts.Token). + SetAPIURL(opts.APIURL). + UseTestServer(opts.UseTestServer). + SetLimiter(limiter). + SetDropRateLimitOverflow(opts.DropRateLimitOverflow). + SetLogFormat(opts.LogFormat). + SetLogFormatter(opts.LogFormatter). + SetHTTPClient(&http.Client{Timeout: httpTimeout}) + api = tgapi.NewAPI(apiOpts) + } uploader := tgapi.NewUploader(api) prefixes := append([]string(nil), opts.Prefixes...) @@ -276,8 +297,13 @@ func (bot *Bot[T]) SetLogger(l *sneklog.Logger) *Bot[T] { } return bot } + if l == bot.logger { + return bot + } bot.addTokenReplacer(l) + bot.closeReplacedLogger(bot.logger, bot.loggerOwned, l, bot.requestLogger, bot.webhookLogger) bot.logger = l + bot.loggerOwned = false return bot } @@ -292,8 +318,13 @@ func (bot *Bot[T]) SetRequestLogger(l *sneklog.Logger) *Bot[T] { } return bot } + if l == bot.requestLogger { + return bot + } bot.addTokenReplacer(l) + bot.closeReplacedLogger(bot.requestLogger, bot.requestLoggerOwned, l, bot.logger, bot.webhookLogger) bot.requestLogger = l + bot.requestLoggerOwned = false return bot } @@ -308,8 +339,13 @@ func (bot *Bot[T]) SetWebhookLogger(l *sneklog.Logger) *Bot[T] { } return bot } + if l == bot.webhookLogger { + return bot + } bot.addTokenReplacer(l) + bot.closeReplacedLogger(bot.webhookLogger, bot.webhookLoggerOwned, l, bot.logger, bot.requestLogger) bot.webhookLogger = l + bot.webhookLoggerOwned = false return bot } @@ -325,8 +361,8 @@ func (bot *Bot[T]) GetUploader() *tgapi.Uploader { return bot.uploader } // - The asynchronous observer dispatcher, after draining queued events // - Registered plugins via Plugin.Close // - Webhook logger (if initialized) -// - Uploader (waits for pending uploads) -// - API client internals +// - Uploader logger resources +// - API client internals, after pending API and upload requests complete // - RequestLogger (if enabled) // - Main logger // @@ -346,18 +382,34 @@ func (bot *Bot[T]) Close() error { } e = append(e, err) } - bot.stopObserverDispatcher() + observerCtx, observerCancel := context.WithTimeout(context.Background(), observerShutdownTimeout) + logCloseErr(bot.stopObserverDispatcher(observerCtx)) + observerCancel() for _, p := range bot.plugins { if err := p.Close(); err != nil { e = append(e, err) } } - if bot.webhookLogger != nil { - if err := bot.webhookLogger.Close(); err != nil { - logCloseErr(err) + closedLoggers := make(map[*sneklog.Logger]struct{}, 3) + closeOwnedLogger := func(logger *sneklog.Logger, owned bool) { + if logger == nil || !owned { + return } + if _, exists := closedLoggers[logger]; exists { + return + } + closedLoggers[logger] = struct{}{} + logCloseErr(logger.Close()) + } + for _, logger := range bot.detachedOwnedLoggers { + closeOwnedLogger(logger, true) + } + bot.detachedOwnedLoggers = nil + if bot.webhookLogger != nil { + closeOwnedLogger(bot.webhookLogger, bot.webhookLoggerOwned) bot.webhookLogger = nil + bot.webhookLoggerOwned = false } if bot.uploader != nil { if err := bot.uploader.Close(); err != nil { @@ -370,14 +422,14 @@ func (bot *Bot[T]) Close() error { } } if bot.requestLogger != nil { - if err := bot.requestLogger.Close(); err != nil { - logCloseErr(err) - } + closeOwnedLogger(bot.requestLogger, bot.requestLoggerOwned) + bot.requestLogger = nil + bot.requestLoggerOwned = false } if bot.logger != nil { - if err := bot.logger.Close(); err != nil { - e = append(e, err) - } + closeOwnedLogger(bot.logger, bot.loggerOwned) + bot.logger = nil + bot.loggerOwned = false } return errors.Join(e...) } @@ -467,25 +519,31 @@ func (bot *Bot[T]) RunWithContext(ctx context.Context) error { defer bot.finishRun() if !bot.useReqLogger && bot.requestLogger != nil { bot.logger.Warnln("Opts#UseRequestLogger is false, but Bot#requestLogger present. Remove Bot#SetRequestLogger or set Opts#UseRequestLogger to true!") - err := bot.requestLogger.Close() - if err != nil { - bot.logger.Errorln(err) + if bot.requestLoggerOwned { + if err := bot.requestLogger.Close(); err != nil { + bot.logger.Errorln(err) + } } bot.requestLogger = nil + bot.requestLoggerOwned = false } if bot.webhookLogger != nil { bot.logger.Warnln("Bot#webhookLogger present. You shouldn't set this, if ran in Long Polling mode!") - err := bot.webhookLogger.Close() - if err != nil { - bot.logger.Errorln(err) + if bot.webhookLoggerOwned { + if err := bot.webhookLogger.Close(); err != nil { + bot.logger.Errorln(err) + } } bot.webhookLogger = nil + bot.webhookLoggerOwned = false } bot.ExecRunners(ctx) // Start update polling in a goroutine + pollDone := make(chan error, 1) go func() { + var terminalErr error defer func() { if r := recover(); r != nil { bot.logger.Errorln(fmt.Sprintf("panic in update polling: %v", r)) @@ -493,6 +551,7 @@ func (bot *Bot[T]) RunWithContext(ctx context.Context) error { if !ok { err = fmt.Errorf("%v", r) } + terminalErr = fmt.Errorf("update polling: %w: %v", ErrHandlerPanic, err) bot.safeEmitEvent(ctx, ErrorEvent{ Plugin: "bot", HandlerKind: HandlerPollingKind, @@ -502,6 +561,7 @@ func (bot *Bot[T]) RunWithContext(ctx context.Context) error { }) } close(bot.updateQueue) + pollDone <- terminalErr }() backoffDelay := time.Duration(0) retryCount := 0 @@ -566,7 +626,8 @@ func (bot *Bot[T]) RunWithContext(ctx context.Context) error { bot.runnerOnceWG.Wait() bot.runnerBgWG.Wait() - return nil + bot.middlewareWG.Wait() + return <-pollDone } // Run starts the bot using a background context. diff --git a/bot_opts.go b/bot_opts.go index 57676c8..49478ea 100644 --- a/bot_opts.go +++ b/bot_opts.go @@ -54,8 +54,8 @@ type BotOpts struct { // Telegram allows up to 30 req/s for most bots. Defaults to 30. RateLimit int - // DropRateLimitOverflow drops incoming updates when rate limit is exceeded instead of queuing. - // Use this to prioritize responsiveness over reliability. + // DropRateLimitOverflow rejects outgoing Telegram API requests immediately when + // rate-limit capacity is unavailable instead of waiting for capacity. DropRateLimitOverflow bool // StrictPayloadType disables callback payload fallback decoding. @@ -96,7 +96,7 @@ type BotOpts struct { // - USE_TEST_SERVER: "true" to use Telegram test server // - API_URL: custom API endpoint // - RATE_LIMIT: max requests per second (default: 30) -// - DROP_RL_OVERFLOW: "true" to drop updates on rate limit overflow +// - DROP_RL_OVERFLOW: "true" to reject rate-limited API requests instead of waiting // - STRICT_PAYLOAD_TYPE: "true" to reject callback payloads encoded in a different format // - MAX_WORKERS: maximum number of concurrent update handlers (default: 32) // - POLL_TIMEOUT: long-polling timeout in seconds for getUpdates (default: 30) @@ -237,8 +237,8 @@ func (opts *BotOpts) SetRateLimit(limit int) *BotOpts { return opts } -// SetDropRateLimitOverflow drops incoming updates when rate limit is exceeded instead of queuing. -// Use this to prioritize responsiveness over reliability. Default is false. +// SetDropRateLimitOverflow configures outgoing Telegram API requests to fail +// immediately when rate-limit capacity is unavailable. Default is false. func (opts *BotOpts) SetDropRateLimitOverflow(drop bool) *BotOpts { opts.DropRateLimitOverflow = drop return opts diff --git a/bot_opts_loader.go b/bot_opts_loader.go index 1818dbc..a4f162a 100644 --- a/bot_opts_loader.go +++ b/bot_opts_loader.go @@ -97,6 +97,9 @@ func (codec BotOptsFileJSONCodec) FromBytes(data []byte) (*BotOpts, error) { // ToBytes encodes BotOpts into JSON file bytes. func (codec BotOptsFileJSONCodec) ToBytes(opts *BotOpts) ([]byte, error) { + if opts == nil { + return nil, ErrOptsIsNil + } fileOpts := &BotOptsFileJSON{ Version: ConfigVersion, Token: opts.Token, @@ -143,7 +146,7 @@ func (codec BotOptsFileJSONCodec) EscapeEnv(s string) string { return string(data[1 : len(data)-1]) } -var envParameterRegex = regexp.MustCompile(`\{\{\s*(\w+)\s*\}\}`) +var envParameterRegex = regexp.MustCompile(`\{\{\s*(\w+)\s*}}`) type botOptsFileEnvEscaper interface { EscapeEnv(string) string @@ -159,6 +162,9 @@ type BotOptsFileCodec interface { // LoadBotOptsFile reads a config file, expands env placeholders, and decodes BotOpts. func LoadBotOptsFile(codec BotOptsFileCodec, filename string) (*BotOpts, error) { + if isNilValue(codec) { + return nil, ErrCodecIsNil + } f, err := os.Open(filename) if err != nil { return nil, err @@ -174,6 +180,12 @@ func LoadBotOptsFile(codec BotOptsFileCodec, filename string) (*BotOpts, error) // SaveBotOptsFile encodes BotOpts with codec and writes the result to filename. func SaveBotOptsFile(codec BotOptsFileCodec, filename string, opts *BotOpts) error { + if isNilValue(codec) { + return ErrCodecIsNil + } + if opts == nil { + return ErrOptsIsNil + } data, err := codec.ToBytes(opts) if err != nil { return err diff --git a/bot_opts_loader_test.go b/bot_opts_loader_test.go index 2c78713..44439e8 100644 --- a/bot_opts_loader_test.go +++ b/bot_opts_loader_test.go @@ -46,6 +46,21 @@ func TestBotOptsFileJSONCodecRoundTrip(t *testing.T) { } } +func TestBotOptsFileRejectsNilInputs(t *testing.T) { + if _, err := (BotOptsFileJSONCodec{}).ToBytes(nil); !errors.Is(err, ErrOptsIsNil) { + t.Fatalf("ToBytes error = %v, want ErrOptsIsNil", err) + } + if _, err := LoadBotOptsFile(nil, "unused"); !errors.Is(err, ErrCodecIsNil) { + t.Fatalf("LoadBotOptsFile error = %v, want ErrCodecIsNil", err) + } + if err := SaveBotOptsFile(nil, "unused", &BotOpts{}); !errors.Is(err, ErrCodecIsNil) { + t.Fatalf("SaveBotOptsFile error = %v, want ErrCodecIsNil", err) + } + if err := SaveBotOptsFile(BotOptsFileJSONCodec{}, "unused", nil); !errors.Is(err, ErrOptsIsNil) { + t.Fatalf("SaveBotOptsFile nil opts error = %v, want ErrOptsIsNil", err) + } +} + func TestLoadBotOptsFileExpandsEnvPlaceholders(t *testing.T) { t.Setenv("TG_TOKEN", "TOKEN_FROM_ENV") t.Setenv("BOT_API_URL", "https://api.example.invalid") diff --git a/bot_test.go b/bot_test.go index ffbebfe..8087237 100644 --- a/bot_test.go +++ b/bot_test.go @@ -22,6 +22,14 @@ func (f pollingRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, erro return f(req) } +type closeCountingWriter struct{ closes int } + +func (w *closeCountingWriter) Close() error { w.closes++; return nil } +func (w *closeCountingWriter) Write(p []byte) (int, error) { return len(p), nil } +func (w *closeCountingWriter) Print(sneklog.LogLevel, string, []*sneklog.MethodTraceback, ...any) error { + return nil +} + type pollingRetryObserver struct { recordingObserver cancel context.CancelFunc @@ -343,6 +351,66 @@ func TestLoggerConfigurationRejectsNilAndLateMutation(t *testing.T) { } } +func TestCloseRespectsLoggerOwnershipAndAliases(t *testing.T) { + ownedWriter := new(closeCountingWriter) + owned := sneklog.NewLogger().AddWriter(ownedWriter) + bot := &Bot[NoData]{ + logger: owned, loggerOwned: true, + requestLogger: owned, requestLoggerOwned: false, + } + if err := bot.Close(); err != nil { + t.Fatalf("Close returned error: %v", err) + } + if ownedWriter.closes != 1 { + t.Fatalf("owned aliased logger closed %d times, want 1", ownedWriter.closes) + } + + callerWriter := new(closeCountingWriter) + callerLogger := sneklog.NewLogger().AddWriter(callerWriter) + bot = &Bot[NoData]{logger: callerLogger, requestLogger: callerLogger, webhookLogger: callerLogger} + if err := bot.Close(); err != nil { + t.Fatalf("Close with caller logger returned error: %v", err) + } + if callerWriter.closes != 0 { + t.Fatalf("caller-owned logger closed %d times", callerWriter.closes) + } + _ = callerLogger.Close() +} + +func TestRunWithContextReturnsPollingPanic(t *testing.T) { + bot := &Bot[NoData]{ + logger: sneklog.NewLogger(), prefixes: []string{"/"}, + plugins: []Plugin[NoData]{{name: "demo"}}, updateQueue: make(chan *tgapi.Update, 1), maxWorkers: 1, + } + defer func() { _ = bot.logger.Close() }() + + err := bot.RunWithContext(context.Background()) + if !errors.Is(err, ErrHandlerPanic) { + t.Fatalf("RunWithContext error = %v, want ErrHandlerPanic", err) + } +} + +func TestNewBotWithAPIUsesInjectedClient(t *testing.T) { + client := &http.Client{Transport: pollingRoundTripFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"ok":true,"result":{"id":7,"is_bot":true,"first_name":"Test","username":"test_bot"}}`)), + }, nil + })} + api := tgapi.NewAPI(tgapi.NewAPIOpts("token").SetAPIURL("http://example.invalid").SetHTTPClient(client)) + bot, err := NewBotWithAPI[NoData](&BotOpts{Token: "token"}, api) + if err != nil { + t.Fatalf("NewBotWithAPI returned error: %v", err) + } + if bot.GetAPI() != api || bot.GetUploader() == nil { + t.Fatal("injected API was not used consistently") + } + if err := bot.Close(); err != nil { + t.Fatalf("Close returned error: %v", err) + } +} + func TestAddAppDataLoggerWriterRejectsNilResult(t *testing.T) { bot := &Bot[NoData]{logger: sneklog.NewLogger(), hasAppData: true} defer func() { _ = bot.logger.Close() }() diff --git a/bot_utils.go b/bot_utils.go index f47dc07..2e14406 100644 --- a/bot_utils.go +++ b/bot_utils.go @@ -18,6 +18,8 @@ import ( "github.com/alitto/pond/v2" ) +const observerShutdownTimeout = 5 * time.Second + func (bot *Bot[T]) addTokenReplacer(loggers ...*sneklog.Logger) { if bot.token == "" { return @@ -30,6 +32,19 @@ func (bot *Bot[T]) addTokenReplacer(loggers ...*sneklog.Logger) { } } +func (bot *Bot[T]) closeReplacedLogger(old *sneklog.Logger, owned bool, replacements ...*sneklog.Logger) { + if old == nil || !owned || len(replacements) == 0 || replacements[0] == old { + return + } + if slices.Contains(replacements[1:], old) { + bot.detachedOwnedLoggers = appendUniqueLogger(bot.detachedOwnedLoggers, old) + return + } + if err := old.Close(); err != nil && bot.logger != nil && bot.logger != old { + bot.logger.Errorln(err) + } +} + func appendUniqueLogger(loggers []*sneklog.Logger, logger *sneklog.Logger) []*sneklog.Logger { if logger == nil { return loggers @@ -80,26 +95,32 @@ func (bot *Bot[T]) initLoggers(opts *BotOpts) { format, formatter := opts.LogFormat, opts.LogFormatter if bot.logger == nil { bot.logger = utils.CreateLogger("BOT", level, format, formatter) + bot.loggerOwned = true if opts.WriteToFile { path := fmt.Sprintf("%s/main.log", strings.TrimRight(opts.LoggerBasePath, "/")) logger, err := utils.CreateFileLogger("BOT", level, path, format, formatter) if err != nil { bot.logger.Errorln(err) } else { + _ = bot.logger.Close() bot.logger = logger + bot.loggerOwned = true } } } if opts.UseRequestLogger && bot.requestLogger == nil { bot.requestLogger = utils.CreateLogger("REQUESTS", level, format, formatter) + bot.requestLoggerOwned = true if opts.WriteToFile { path := fmt.Sprintf("%s/requests.log", strings.TrimRight(opts.LoggerBasePath, "/")) logger, err := utils.CreateFileLogger("REQUESTS", level, path, format, formatter) if err != nil { bot.logger.Errorln(err) } else { + _ = bot.requestLogger.Close() bot.requestLogger = logger + bot.requestLoggerOwned = true } } } @@ -124,17 +145,29 @@ func (bot *Bot[T]) beginRun() error { } func (bot *Bot[T]) finishRun() { - bot.stopObserverDispatcher() + ctx, cancel := context.WithTimeout(context.Background(), observerShutdownTimeout) + if err := bot.stopObserverDispatcher(ctx); err != nil && bot.logger != nil { + bot.logger.Errorln(err) + } + cancel() bot.runStateMu.Lock() bot.running = false bot.runStateMu.Unlock() } -func (bot *Bot[T]) stopObserverDispatcher() { +func (bot *Bot[T]) stopObserverDispatcher(ctx context.Context) error { if bot.observerAsync == nil { - return + return nil } - bot.observerAsync.close() + return bot.observerAsync.close(ctx) +} + +func (bot *Bot[T]) startAsyncTask(task func()) { + bot.middlewareWG.Add(1) + go func() { + defer bot.middlewareWG.Done() + task() + }() } func nextPollRetryDelay(prev time.Duration) time.Duration { diff --git a/bot_webhook.go b/bot_webhook.go index 82ec7e4..e6be725 100644 --- a/bot_webhook.go +++ b/bot_webhook.go @@ -257,10 +257,13 @@ func (bot *Bot[T]) CloseWebhook() error { } } if bot.webhookLogger != nil { - if err := bot.webhookLogger.Close(); err != nil { - e = append(e, err) + if bot.webhookLoggerOwned { + if err := bot.webhookLogger.Close(); err != nil { + e = append(e, err) + } } bot.webhookLogger = nil + bot.webhookLoggerOwned = false } return errors.Join(e...) } @@ -283,6 +286,7 @@ func (bot *Bot[T]) runWebhookRuntime(ctx context.Context, run func(context.Conte if bot.webhookLogger == nil { bot.webhookLogger = utils.CreateLogger("WEBHOOK", bot.GetLoggerLevel(), bot.logFormat, bot.logFormatter) + bot.webhookLoggerOwned = true } bot.addTokenReplacer(bot.webhookLogger) bot.ExecRunners(runCtx) @@ -297,6 +301,7 @@ func (bot *Bot[T]) runWebhookRuntime(ctx context.Context, run func(context.Conte cancel() close(bot.updateQueue) <-workersDone + bot.middlewareWG.Wait() bot.runnerOnceWG.Wait() bot.runnerBgWG.Wait() diff --git a/drafts.go b/drafts.go index 13a7ac6..e6ab005 100644 --- a/drafts.go +++ b/drafts.go @@ -45,10 +45,8 @@ type DraftProvider struct { generator draftIDGenerator } -// NewRandomDraftProvider creates a new DraftProvider using random draft IDs. -// -// The provider will use random numbers for draft IDs. -// All drafts created via this provider will have unpredictable, unique IDs. +// NewRandomDraftProvider creates a DraftProvider using non-cryptographic random IDs. +// Zero values and collisions with active drafts are retried. func NewRandomDraftProvider(api *tgapi.API) *DraftProvider { return &DraftProvider{ api: api, generator: &RandomDraftIDGenerator{}, diff --git a/errors.go b/errors.go index 5f420cf..85d6feb 100644 --- a/errors.go +++ b/errors.go @@ -38,6 +38,8 @@ var ( ErrAPIIsNil = errors.New("api is nil") // ErrMessageIDZero reports that an operation requires a non-zero message ID. ErrMessageIDZero = errors.New("message ID is zero") + // ErrCodecIsNil reports that a config operation received a nil codec. + ErrCodecIsNil = errors.New("codec is nil") ) var ( // ErrBindArgsTargetNotPointer reports that BindArgs received a nil or non-pointer destination. @@ -66,6 +68,8 @@ var ( ErrHandlerExecutorNil = errors.New("handler executor is nil") // ErrHandlerPanic reports a panic recovered from a user handler. ErrHandlerPanic = errors.New("handler panicked") + // ErrObserverShutdownTimeout reports that observer callbacks did not stop before shutdown timed out. + ErrObserverShutdownTimeout = errors.New("observer shutdown timed out") // ErrInlineKeyboardButtonAction reports a button without exactly one action. ErrInlineKeyboardButtonAction = errors.New("inline keyboard button must have exactly one action") // ErrCallbackDataLength reports callback data outside Telegram's 1-64 byte range. diff --git a/handler.go b/handler.go index fc3cf24..b27ee99 100644 --- a/handler.go +++ b/handler.go @@ -52,13 +52,12 @@ func (bot *Bot[T]) handle(parentCtx context.Context, u *tgapi.Update) { sceneRuntime: bot, observer: bot.observer, eventEmitter: bot.safeEmitEvent, + asyncTask: bot.startAsyncTask, payloadType: bot.payloadType, botID: bot.userID, ctx: ctx, } bot.prepareUpdateCtx(u, msgCtx) - unlockScenes := bot.sceneLocks.lock(sceneKeysForContext(msgCtx)) - defer unlockScenes() bot.safeEmitEvent(ctx, UpdateReceivedEvent{ UpdateID: u.UpdateID, UpdateType: u.Type, diff --git a/handler_test.go b/handler_test.go index 70e99cf..e4fc34a 100644 --- a/handler_test.go +++ b/handler_test.go @@ -7,6 +7,7 @@ import ( "net/http" "strings" "testing" + "time" "git.scuroneko.dev/scuroneko/laniakea/tgapi" "git.scuroneko.dev/scuroneko/sneklog/v2" @@ -1301,6 +1302,8 @@ func TestParseCommandTable(t *testing.T) { {name: "plain text", text: "hello", wantPrefix: "", wantCmd: "", wantArgs: ""}, {name: "command no args", text: "/start", wantPrefix: "/", wantCmd: "start", wantArgs: ""}, {name: "command with args", text: "/ban 42 reason", wantPrefix: "/", wantCmd: "ban", wantArgs: "42 reason"}, + {name: "tab separator", text: "/ban\t42", wantPrefix: "/", wantCmd: "ban", wantArgs: "42"}, + {name: "newline separator", text: "/ban\n42", wantPrefix: "/", wantCmd: "ban", wantArgs: "42"}, {name: "alternate prefix", text: "!ping", wantPrefix: "!", wantCmd: "ping", wantArgs: ""}, {name: "leading space after prefix", text: "/ start now", wantPrefix: "/", wantCmd: "start", wantArgs: "now"}, {name: "command with botname", text: "/start@mybot extra", wantPrefix: "/", wantCmd: "start@mybot", wantArgs: "extra"}, @@ -1332,6 +1335,7 @@ func TestHandleMessageStripsBotUsernameSuffix(t *testing.T) { }{ {name: "matching botname", botUsername: "mybot", text: "/start@mybot hello", wantCalled: true}, {name: "matching botname no args", botUsername: "mybot", text: "/start@mybot", wantCalled: true}, + {name: "matching botname ignores case", botUsername: "MyBot", text: "/start@mybot", wantCalled: true}, {name: "other botname", botUsername: "mybot", text: "/start@otherbot hello", wantCalled: false}, {name: "no botname", botUsername: "mybot", text: "/start hello", wantCalled: true}, {name: "bot has no username", botUsername: "", text: "/start@mybot hello", wantCalled: false}, @@ -1581,6 +1585,75 @@ func sceneMessageUpdate(updateID int, text string) tgapi.Update { } } +func TestHandleDoesNotSerializeUnrelatedUsersWithoutScene(t *testing.T) { + started := make(chan struct{}, 2) + release := make(chan struct{}) + plugin := NewPlugin[NoData]("commands") + plugin.Command("work", func(*MessageContext, NoData) error { + started <- struct{}{} + <-release + return nil + }) + bot := &Bot[NoData]{ + logger: sneklog.NewLogger(), + plugins: []Plugin[NoData]{clonePlugin(plugin)}, + prefixes: []string{"/"}, + sessionStore: NewMemorySessionStore(), + sceneScopePriority: []SceneScope{SceneScopeUserChat, SceneScopeChat, SceneScopeUser}, + } + defer func() { _ = bot.logger.Close() }() + + updateFor := func(id int) tgapi.Update { + return tgapi.Update{UpdateID: id, Type: tgapi.UpdateTypeMessage, Message: &tgapi.Message{ + MessageID: id, Text: "/work", From: &tgapi.User{ID: int64(id)}, + Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypeSupergroup}, + }} + } + done := make(chan struct{}, 2) + for id := 1; id <= 2; id++ { + update := updateFor(id) + go func() { + bot.handle(context.Background(), &update) + done <- struct{}{} + }() + } + for range 2 { + select { + case <-started: + case <-time.After(time.Second): + close(release) + t.Fatal("updates from unrelated users in one chat were serialized") + } + } + close(release) + <-done + <-done +} + +func TestCommandMiddlewareBlockIsNotReportedAsError(t *testing.T) { + called := false + plugin := NewPlugin[NoData]("commands") + plugin.Command("blocked", func(*MessageContext, NoData) error { + called = true + return nil + }).Use(NewMiddleware("deny", func(*MessageContext, NoData) bool { return false })) + observer := &recordingObserver{} + bot := &Bot[NoData]{ + logger: sneklog.NewLogger(), plugins: []Plugin[NoData]{clonePlugin(plugin)}, + prefixes: []string{"/"}, observer: observer, sessionStore: NewMemorySessionStore(), + } + defer func() { _ = bot.logger.Close() }() + update := sceneMessageUpdate(201, "/blocked") + bot.handle(context.Background(), &update) + + if called { + t.Fatal("command executed after middleware blocked it") + } + if len(observer.errors) != 0 { + t.Fatalf("middleware block emitted errors: %#v", observer.errors) + } +} + func TestHandleCallbackObserverEmitsDecodeErrors(t *testing.T) { observer := &recordingObserver{} bot := &Bot[NoData]{ diff --git a/msg_context.go b/msg_context.go index a1be0fc..e1e93e5 100644 --- a/msg_context.go +++ b/msg_context.go @@ -84,45 +84,59 @@ type MessageContext struct { sceneRuntime sceneRuntime observer Observer eventEmitter func(context.Context, Event) + asyncTask func(func()) botID int64 ctx context.Context } +func (ctx *MessageContext) buildEditMessageTextParams(messageID int, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) (tgapi.EditMessageText, error) { + params := tgapi.EditMessageText{} + switch { + case messageID > 0 && ctx.Msg != nil && ctx.Msg.Chat != nil: + params.MessageID = messageID + params.ChatID = ctx.Msg.Chat.ID + params.BusinessConnectionID = ctx.Msg.BusinessConnectionID + case ctx.InlineMsgID != "": + params.InlineMessageID = ctx.InlineMsgID + default: + return params, ErrEditTargetMissing + } + if parseMode != "" { + params.ParseMode = parseMode + } + if keyboard != nil { + params.ReplyMarkup = keyboard.Get() + } + return params, nil +} + // AnswerMessage represents a message sent or edited via MessageContext. // It holds metadata to allow further editing or deletion. type AnswerMessage struct { // MessageID identifies the sent Telegram message. MessageID int // Text contains the text or caption sent with the message. + // For rich messages, it contains rendered HTML for v1 compatibility. Text string + // RichHTML contains the rendered HTML of a rich message. + RichHTML string // Since: Bot API 10.2 // IsMedia reports whether the answer contains media. IsMedia bool ctx *MessageContext // internal back-reference } func (ctx *MessageContext) edit(messageID int, text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { + params, err := ctx.buildEditMessageTextParams(messageID, keyboard, parseMode) + if err != nil { + ctx.Logger.Errorln(err) + return nil + } if err := validateMessageText(text); err != nil { ctx.Logger.Errorln(err) return nil } - params := tgapi.EditMessageText{ - Text: text, - ParseMode: parseMode, - } - switch { - case messageID > 0 && ctx.Msg != nil: - params.MessageID = messageID - params.ChatID = ctx.Msg.Chat.ID - case ctx.InlineMsgID != "": - params.InlineMessageID = ctx.InlineMsgID - default: - ctx.Logger.Errorln(ErrEditTargetMissing) - return nil - } - if keyboard != nil { - params.ReplyMarkup = keyboard.Get() - } + params.Text = text msg, _, err := ctx.API.EditMessageTextWithContext(ctx.Context(), params) if err != nil { ctx.Logger.Errorln(err) @@ -193,9 +207,10 @@ func (ctx *MessageContext) editPhotoText(messageID int, text string, kb *InlineK ParseMode: parseMode, } switch { - case messageID > 0 && ctx.Msg != nil: + case messageID > 0 && ctx.Msg != nil && ctx.Msg.Chat != nil: params.ChatID = ctx.Msg.Chat.ID params.MessageID = messageID + params.BusinessConnectionID = ctx.Msg.BusinessConnectionID case ctx.InlineMsgID != "": params.InlineMessageID = ctx.InlineMsgID default: @@ -245,7 +260,7 @@ func (m *AnswerMessage) EditCaptionKeyboardMarkdown(text string, kb *InlineKeybo } func (ctx *MessageContext) answer(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { - if ctx.Msg == nil { + if ctx.Msg == nil || ctx.Msg.Chat == nil { ctx.Logger.Errorln(ErrMessageContextNil) return nil } @@ -254,9 +269,10 @@ func (ctx *MessageContext) answer(text string, keyboard *InlineKeyboard, parseMo return nil } params := tgapi.SendMessage{ - ChatID: ctx.Msg.Chat.ID, - Text: text, - ParseMode: parseMode, + BusinessConnectionID: ctx.Msg.BusinessConnectionID, + ChatID: ctx.Msg.Chat.ID, + Text: text, + ParseMode: parseMode, } if keyboard != nil { params.ReplyMarkup = keyboard.Get() @@ -339,7 +355,7 @@ func (ctx *MessageContext) answerLong(text string, keyboard *InlineKeyboard, par ctx.Logger.Errorln(ErrMessageSplitImpossible) return nil } - if ctx.Msg == nil { + if ctx.Msg == nil || ctx.Msg.Chat == nil { ctx.Logger.Errorln(ErrMessageContextNil) return nil } @@ -374,7 +390,7 @@ func (ctx *MessageContext) answerLong(text string, keyboard *InlineKeyboard, par } func (ctx *MessageContext) answerPhoto(photoID, text string, kb *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { - if ctx.Msg == nil { + if ctx.Msg == nil || ctx.Msg.Chat == nil { ctx.Logger.Errorln(ErrMessageContextNil) return nil } @@ -383,10 +399,11 @@ func (ctx *MessageContext) answerPhoto(photoID, text string, kb *InlineKeyboard, return nil } params := tgapi.SendPhoto{ - ChatID: ctx.Msg.Chat.ID, - Caption: text, - ParseMode: parseMode, - Photo: photoID, + BusinessConnectionID: ctx.Msg.BusinessConnectionID, + ChatID: ctx.Msg.Chat.ID, + Caption: text, + ParseMode: parseMode, + Photo: photoID, } if kb != nil { params.ReplyMarkup = kb.Get() @@ -449,7 +466,7 @@ func (ctx *MessageContext) delete(messageID int) { ctx.Logger.Errorln(ErrMessageIDZero) return } - if ctx.Msg == nil { + if ctx.Msg == nil || ctx.Msg.Chat == nil { ctx.Logger.Errorln(ErrMessageContextNil) return } @@ -501,12 +518,14 @@ func (ctx *MessageContext) AnswerCallbackURL(u string) { ctx.answerCallbackQuery // SendAction sends a chat action (typing, uploading_photo, etc.) to indicate bot activity. func (ctx *MessageContext) SendAction(action tgapi.ChatActionType) { - if ctx.Msg == nil { + if ctx.Msg == nil || ctx.Msg.Chat == nil { ctx.Logger.Errorln(ErrMessageContextNil) return } params := tgapi.SendChatAction{ - ChatID: ctx.Msg.Chat.ID, Action: action, + BusinessConnectionID: ctx.Msg.BusinessConnectionID, + ChatID: ctx.Msg.Chat.ID, + Action: action, } if ctx.Msg.MessageThreadID > 0 { params.MessageThreadID = ctx.Msg.MessageThreadID @@ -543,7 +562,7 @@ func (ctx *MessageContext) error(err error) { func (ctx *MessageContext) Error(err error) { ctx.error(err) } func (ctx *MessageContext) newDraft(parseMode tgapi.ParseMode) *Draft { - if ctx.Msg == nil { + if ctx.Msg == nil || ctx.Msg.Chat == nil { ctx.Logger.Errorln(ErrMessageContextNil) return nil } @@ -827,7 +846,7 @@ func (ctx *MessageContext) UpsertKeyboardMarkdown(text string, keyboard *InlineK } func (ctx *MessageContext) richAnswer(rich tgapi.InputRichMessage, keyboard *InlineKeyboard) *AnswerMessage { - if ctx.Msg == nil { + if ctx.Msg == nil || ctx.Msg.Chat == nil { ctx.Logger.Errorln(ErrMessageContextNil) return nil } @@ -844,6 +863,9 @@ func (ctx *MessageContext) richAnswer(rich tgapi.InputRichMessage, keyboard *Inl if ctx.Msg.DirectMessageTopic != nil { params.DirectMessagesTopicID = ctx.Msg.DirectMessageTopic.TopicID } + if ctx.Msg.BusinessConnectionID != "" { + params.BusinessConnectionID = ctx.Msg.BusinessConnectionID + } msg, err := ctx.API.SendRichMessageWithContext(ctx.Context(), params) if err != nil { @@ -851,7 +873,11 @@ func (ctx *MessageContext) richAnswer(rich tgapi.InputRichMessage, keyboard *Inl return nil } return &AnswerMessage{ - MessageID: msg.MessageID, ctx: ctx, Text: rich.HTML, IsMedia: false, + MessageID: msg.MessageID, + Text: rich.HTML, + RichHTML: rich.HTML, + IsMedia: false, + ctx: ctx, } } @@ -877,3 +903,84 @@ func (ctx *MessageContext) RichAnswer(blocks ...tgapi.InputRichBlock) *AnswerMes func (ctx *MessageContext) RichAnswerKeyboard(keyboard *InlineKeyboard, blocks ...tgapi.InputRichBlock) *AnswerMessage { return ctx.richBlocksAnswer(keyboard, blocks...) } + +func (ctx *MessageContext) editRich(messageID int, keyboard *InlineKeyboard, blocks ...tgapi.InputRichBlock) *AnswerMessage { + params, err := ctx.buildEditMessageTextParams(messageID, keyboard, "") + if err != nil { + ctx.Logger.Errorln(err) + return nil + } + rich, err := tgrich.BuildHTML(blocks...) + if err != nil { + ctx.Logger.Errorln(err) + return nil + } + params.RichMessage = &rich + msg, _, err := ctx.API.EditMessageTextWithContext(ctx.Context(), params) + if err != nil { + ctx.Logger.Errorln(err) + return nil + } + resultMessageID := messageID + if msg.MessageID > 0 { + resultMessageID = msg.MessageID + } + return &AnswerMessage{ + MessageID: resultMessageID, Text: rich.HTML, RichHTML: rich.HTML, IsMedia: false, ctx: ctx, + } +} +func (ctx *MessageContext) editRichCallback(keyboard *InlineKeyboard, blocks ...tgapi.InputRichBlock) *AnswerMessage { + if ctx.CallbackMsgID == 0 && ctx.InlineMsgID == "" { + ctx.Logger.Errorln(ErrCallbackMessageMissing) + return nil + } + return ctx.editRich(ctx.CallbackMsgID, keyboard, blocks...) +} +func (ctx *MessageContext) upsertRichKeyboard(keyboard *InlineKeyboard, blocks ...tgapi.InputRichBlock) *AnswerMessage { + if ctx.IsCallback() { + if ctx.HasPhoto() { + rich, err := tgrich.BuildHTML(blocks...) + if err != nil { + ctx.Logger.Errorln(err) + return nil + } + ctx.CallbackDelete() + return ctx.richAnswer(rich, keyboard) + } + return ctx.editRichCallback(keyboard, blocks...) + } + return ctx.richBlocksAnswer(keyboard, blocks...) +} + +// EditCallbackRich builds rich blocks and replaces the callback message content and inline keyboard. +// It doesn't upload local files referenced with attach://. +// +// Since: Bot API 10.2 +func (ctx *MessageContext) EditCallbackRich(keyboard *InlineKeyboard, blocks ...tgapi.InputRichBlock) *AnswerMessage { + return ctx.editRichCallback(keyboard, blocks...) +} + +// UpsertKeyboardRich builds rich blocks and either edits the callback message or sends a new message. +// Photo callback messages are replaced because their text content can't be edited directly. +// It doesn't upload local files referenced with attach://. +// +// Since: Bot API 10.2 +func (ctx *MessageContext) UpsertKeyboardRich(keyboard *InlineKeyboard, blocks ...tgapi.InputRichBlock) *AnswerMessage { + return ctx.upsertRichKeyboard(keyboard, blocks...) +} + +// EditRich builds rich blocks and replaces the message content without changing its inline keyboard. +// It doesn't upload local files referenced with attach://. +// +// Since: Bot API 10.2 +func (m *AnswerMessage) EditRich(blocks ...tgapi.InputRichBlock) *AnswerMessage { + return m.ctx.editRich(m.MessageID, nil, blocks...) +} + +// EditRichKeyboard builds rich blocks and replaces the message content and inline keyboard. +// It doesn't upload local files referenced with attach://. +// +// Since: Bot API 10.2 +func (m *AnswerMessage) EditRichKeyboard(keyboard *InlineKeyboard, blocks ...tgapi.InputRichBlock) *AnswerMessage { + return m.ctx.editRich(m.MessageID, keyboard, blocks...) +} diff --git a/msg_context_test.go b/msg_context_test.go index 7700ad3..a862531 100644 --- a/msg_context_test.go +++ b/msg_context_test.go @@ -14,6 +14,101 @@ import ( "git.scuroneko.dev/scuroneko/sneklog/v2" ) +func newMessageContextTestAPI(t *testing.T, transport roundTripFunc) *tgapi.API { + t.Helper() + api := tgapi.NewAPI( + tgapi.NewAPIOpts("token"). + SetAPIURL("https://example.test"). + SetHTTPClient(&http.Client{Transport: transport}), + ) + t.Cleanup(func() { + if err := api.Close(); err != nil { + t.Fatalf("Close returned error: %v", err) + } + }) + return api +} + +func readMessageContextRequest(t *testing.T, req *http.Request) map[string]any { + t.Helper() + body, err := io.ReadAll(req.Body) + if err != nil { + t.Fatalf("failed to read request body: %v", err) + } + var decoded map[string]any + if err := json.Unmarshal(body, &decoded); err != nil { + t.Fatalf("failed to decode request body: %v", err) + } + return decoded +} + +func messageContextResponse(result string) *http.Response { + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"ok":true,"result":` + result + `}`)), + } +} + +func TestMessageContextPropagatesBusinessConnection(t *testing.T) { + tests := []struct { + name string + wantMethod string + result string + invoke func(*MessageContext) + }{ + {name: "send message", wantMethod: "sendMessage", result: `{"message_id":11,"date":1}`, invoke: func(ctx *MessageContext) { ctx.Answer("text") }}, + {name: "send photo", wantMethod: "sendPhoto", result: `{"message_id":11,"date":1}`, invoke: func(ctx *MessageContext) { ctx.AnswerPhoto("photo-id", "caption") }}, + {name: "edit caption", wantMethod: "editMessageCaption", result: `{"message_id":11,"date":1}`, invoke: func(ctx *MessageContext) { (&AnswerMessage{MessageID: 7, ctx: ctx}).EditCaption("caption") }}, + {name: "send action", wantMethod: "sendChatAction", result: `true`, invoke: func(ctx *MessageContext) { ctx.SendAction(tgapi.ChatActionTyping) }}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var gotBody map[string]any + api := newMessageContextTestAPI(t, roundTripFunc(func(req *http.Request) (*http.Response, error) { + if !strings.HasSuffix(req.URL.Path, "/"+tt.wantMethod) { + t.Fatalf("request path = %q, want method %q", req.URL.Path, tt.wantMethod) + } + gotBody = readMessageContextRequest(t, req) + return messageContextResponse(tt.result), nil + })) + ctx := &MessageContext{ + API: api, + Msg: &tgapi.Message{ + BusinessConnectionID: "business-1", + Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, + }, + Logger: sneklog.NewLogger(), + } + + tt.invoke(ctx) + if got := gotBody["business_connection_id"]; got != "business-1" { + t.Fatalf("business_connection_id = %v, want business-1", got) + } + }) + } +} + +func TestMessageContextHelpersRejectMissingChat(t *testing.T) { + requests := 0 + api := newMessageContextTestAPI(t, roundTripFunc(func(req *http.Request) (*http.Response, error) { + requests++ + return nil, errors.New("unexpected request") + })) + ctx := &MessageContext{API: api, Msg: &tgapi.Message{}, CallbackMsgID: 7, Logger: sneklog.NewLogger()} + + if answer := ctx.Answer("text"); answer != nil { + t.Fatalf("Answer returned %#v for a message without a chat", answer) + } + if answer := ctx.EditCallback("text", nil); answer != nil { + t.Fatalf("EditCallback returned %#v for a message without a chat", answer) + } + if requests != 0 { + t.Fatalf("missing-chat helpers made %d requests", requests) + } +} + func TestRichAnswerBuildsInputBlocks(t *testing.T) { var gotBody map[string]any client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { @@ -33,8 +128,11 @@ func TestRichAnswerBuildsInputBlocks(t *testing.T) { api := tgapi.NewAPI(tgapi.NewAPIOpts("token").SetAPIURL("https://example.test").SetHTTPClient(client)) defer func() { _ = api.Close() }() ctx := &MessageContext{ - API: api, - Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}}, + API: api, + Msg: &tgapi.Message{ + BusinessConnectionID: "business-1", + Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, + }, Logger: sneklog.NewLogger(), } @@ -49,6 +147,12 @@ func TestRichAnswerBuildsInputBlocks(t *testing.T) { if _, exists := rich["skip_entity_detection"]; exists { t.Fatalf("rich_message unexpectedly disables entity detection: %#v", rich) } + if got := gotBody["business_connection_id"]; got != "business-1" { + t.Fatalf("business_connection_id = %v, want business-1", got) + } + if answer.Text != "
ready
" || answer.RichHTML != answer.Text { + t.Fatalf("unexpected answer content: Text=%q RichHTML=%q", answer.Text, answer.RichHTML) + } } func TestRichAnswerRejectsInvalidBlocksWithoutRequest(t *testing.T) { @@ -69,6 +173,190 @@ func TestRichAnswerRejectsInvalidBlocksWithoutRequest(t *testing.T) { } } +func TestAnswerMessageEditRich(t *testing.T) { + tests := []struct { + name string + withKeyboard bool + }{ + {name: "content only"}, + {name: "content and keyboard", withKeyboard: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var gotPath string + var gotBody map[string]any + api := newMessageContextTestAPI(t, roundTripFunc(func(req *http.Request) (*http.Response, error) { + gotPath = req.URL.Path + gotBody = readMessageContextRequest(t, req) + return messageContextResponse(`{"message_id":11,"date":1}`), nil + })) + ctx := &MessageContext{ + API: api, + Msg: &tgapi.Message{ + BusinessConnectionID: "business-1", + Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, + }, + Logger: sneklog.NewLogger(), + } + original := &AnswerMessage{MessageID: 7, ctx: ctx} + block := tgrich.P(tgrich.Bold(tgrich.Text("updated"))) + + var answer *AnswerMessage + if tt.withKeyboard { + kb := NewInlineKeyboardJSON(1).AddCallbackButton("A", "cmd") + answer = original.EditRichKeyboard(kb, block) + } else { + answer = original.EditRich(block) + } + + if answer == nil { + t.Fatal("rich edit returned nil") + } + if gotPath != "/bottoken/editMessageText" { + t.Fatalf("unexpected request path: %s", gotPath) + } + if got := gotBody["chat_id"]; got != float64(42) { + t.Fatalf("chat_id = %v, want 42", got) + } + if got := gotBody["message_id"]; got != float64(7) { + t.Fatalf("message_id = %v, want 7", got) + } + if got := gotBody["business_connection_id"]; got != "business-1" { + t.Fatalf("business_connection_id = %v, want business-1", got) + } + rich, ok := gotBody["rich_message"].(map[string]any) + if !ok || rich["html"] != "updated
" { + t.Fatalf("rich_message = %#v", gotBody["rich_message"]) + } + if _, exists := gotBody["text"]; exists { + t.Fatalf("edit request unexpectedly contains text: %#v", gotBody) + } + _, hasKeyboard := gotBody["reply_markup"] + if hasKeyboard != tt.withKeyboard { + t.Fatalf("reply_markup presence = %v, want %v", hasKeyboard, tt.withKeyboard) + } + if answer.MessageID != 11 || answer.Text != "updated
" || answer.RichHTML != answer.Text { + t.Fatalf("unexpected answer: %#v", answer) + } + if answer.ctx != ctx { + t.Fatal("edited answer lost its message context") + } + }) + } +} + +func TestEditCallbackRichEditsInlineMessage(t *testing.T) { + var gotBody map[string]any + api := newMessageContextTestAPI(t, roundTripFunc(func(req *http.Request) (*http.Response, error) { + gotBody = readMessageContextRequest(t, req) + return messageContextResponse("true"), nil + })) + ctx := &MessageContext{ + API: api, + InlineMsgID: "inline-1", + Logger: sneklog.NewLogger(), + } + kb := NewInlineKeyboardJSON(1).AddCallbackButton("A", "cmd") + + answer := ctx.EditCallbackRich(kb, tgrich.P(tgrich.Text("inline"))) + if answer == nil { + t.Fatal("EditCallbackRich returned nil") + } + if got := gotBody["inline_message_id"]; got != "inline-1" { + t.Fatalf("inline_message_id = %v, want inline-1", got) + } + if _, exists := gotBody["chat_id"]; exists { + t.Fatalf("inline edit unexpectedly contains chat_id: %#v", gotBody) + } + if _, exists := gotBody["business_connection_id"]; exists { + t.Fatalf("inline edit unexpectedly contains business_connection_id: %#v", gotBody) + } + rich, ok := gotBody["rich_message"].(map[string]any) + if !ok || rich["html"] != "inline
" { + t.Fatalf("rich_message = %#v", gotBody["rich_message"]) + } + if _, exists := gotBody["reply_markup"]; !exists { + t.Fatal("inline rich edit has no reply_markup") + } + if answer.MessageID != 0 || answer.Text != "inline
" || answer.RichHTML != answer.Text { + t.Fatalf("unexpected inline answer: %#v", answer) + } +} + +func TestUpsertKeyboardRichValidatesPhotoBeforeDelete(t *testing.T) { + requests := 0 + api := newMessageContextTestAPI(t, roundTripFunc(func(req *http.Request) (*http.Response, error) { + requests++ + t.Fatalf("unexpected request to %s", req.URL.Path) + return nil, nil + })) + ctx := &MessageContext{ + API: api, + CallbackMsgID: 7, + Msg: &tgapi.Message{ + Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, + Photo: []tgapi.PhotoSize{{FileID: "photo-1"}}, + }, + Logger: sneklog.NewLogger(), + } + + answer := ctx.UpsertKeyboardRich(nil, tgrich.H(tgrich.Text("invalid"), 0)) + if answer != nil { + t.Fatalf("UpsertKeyboardRich returned an answer for invalid blocks: %#v", answer) + } + if requests != 0 { + t.Fatalf("invalid photo upsert made %d requests", requests) + } +} + +func TestUpsertKeyboardRichReplacesPhotoCallback(t *testing.T) { + var paths []string + var sendBody map[string]any + api := newMessageContextTestAPI(t, roundTripFunc(func(req *http.Request) (*http.Response, error) { + paths = append(paths, req.URL.Path) + switch req.URL.Path { + case "/bottoken/deleteMessage": + return messageContextResponse("true"), nil + case "/bottoken/sendRichMessage": + sendBody = readMessageContextRequest(t, req) + return messageContextResponse(`{"message_id":12,"date":1}`), nil + default: + t.Fatalf("unexpected request path: %s", req.URL.Path) + return nil, nil + } + })) + ctx := &MessageContext{ + API: api, + CallbackMsgID: 7, + Msg: &tgapi.Message{ + Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, + Photo: []tgapi.PhotoSize{{FileID: "photo-1"}}, + }, + Logger: sneklog.NewLogger(), + } + kb := NewInlineKeyboardJSON(1).AddCallbackButton("A", "cmd") + + answer := ctx.UpsertKeyboardRich(kb, tgrich.P(tgrich.Text("replacement"))) + if answer == nil { + t.Fatal("UpsertKeyboardRich returned nil") + } + wantPaths := []string{"/bottoken/deleteMessage", "/bottoken/sendRichMessage"} + if !reflect.DeepEqual(paths, wantPaths) { + t.Fatalf("request paths = %#v, want %#v", paths, wantPaths) + } + rich, ok := sendBody["rich_message"].(map[string]any) + if !ok || rich["html"] != "replacement
" { + t.Fatalf("rich_message = %#v", sendBody["rich_message"]) + } + if _, exists := sendBody["reply_markup"]; !exists { + t.Fatal("replacement rich message has no reply_markup") + } + if answer.MessageID != 12 || answer.Text != "replacement
" || answer.RichHTML != answer.Text { + t.Fatalf("unexpected replacement answer: %#v", answer) + } +} + func TestAnswerPhotoIncludesDirectMessagesTopicID(t *testing.T) { var gotBody map[string]any diff --git a/msg_handler.go b/msg_handler.go index d3adfa5..d723a59 100644 --- a/msg_handler.go +++ b/msg_handler.go @@ -1,8 +1,10 @@ package laniakea import ( + "errors" "strings" "time" + "unicode" "git.scuroneko.dev/scuroneko/laniakea/tgapi" ) @@ -21,8 +23,9 @@ func (bot *Bot[T]) handleMessage(update *tgapi.Update, ctx *MessageContext) bool if strings.Contains(cmd, "@") { botUsername := bot.username - if botUsername != "" && strings.HasSuffix(cmd, "@"+botUsername) { - cmd = cmd[:len(cmd)-len("@"+botUsername)] // remove @botname + at := strings.LastIndexByte(cmd, '@') + if at > 0 && botUsername != "" && strings.EqualFold(cmd[at+1:], botUsername) { + cmd = cmd[:at] // remove @botname } } @@ -52,6 +55,9 @@ func (bot *Bot[T]) handleMessage(update *tgapi.Update, ctx *MessageContext) bool }) err := plugin.executeCmd(cmd, ctx, bot.appData) + if errors.Is(err, errMiddlewareBlocked) { + err = nil + } handlerEndEvent := HandlerFinishedEvent{ UpdateID: update.UpdateID, UpdateType: update.Type, @@ -227,6 +233,9 @@ func (bot *Bot[T]) handleCallback(update *tgapi.Update, ctx *MessageContext) boo ChatID: ctx.ChatID, }) err := plugin.executePayload(data.Command, ctx, bot.appData) + if errors.Is(err, errMiddlewareBlocked) { + err = nil + } endEvent := HandlerFinishedEvent{ UpdateID: update.UpdateID, @@ -282,7 +291,7 @@ func (bot *Bot[T]) checkPrefixes(text string) (string, bool) { func (bot *Bot[T]) parseCommand(text string) (prefix, cmd, args string) { if prefix, hasPrefix := bot.checkPrefixes(text); hasPrefix { text = strings.TrimSpace(text[len(prefix):]) - spaceIndex := strings.Index(text, " ") + spaceIndex := strings.IndexFunc(text, unicode.IsSpace) var cmd string var args string if spaceIndex == -1 { diff --git a/observer.go b/observer.go index a49ddac..34bf4ab 100644 --- a/observer.go +++ b/observer.go @@ -228,7 +228,8 @@ func (ErrorEvent) isEvent() {} // During RunWithContext and RunWebhookWithContext, callbacks execute on a // dedicated dispatcher goroutine in enqueue order and never block update // handlers. The queue is bounded; overload drops events and emits sampled -// warnings. Runtime shutdown drains events that were already queued. +// warnings. Runtime shutdown cancels callback contexts and drains queued events; +// Bot.Close returns ErrObserverShutdownTimeout if a callback ignores cancellation. type Observer interface { OnUpdateReceived(ctx context.Context, event UpdateReceivedEvent) OnUpdateHandled(ctx context.Context, event UpdateHandledEvent) diff --git a/observer_async.go b/observer_async.go index ef9fff6..071b4ba 100644 --- a/observer_async.go +++ b/observer_async.go @@ -5,6 +5,7 @@ import ( "fmt" "sync" "sync/atomic" + "time" "git.scuroneko.dev/scuroneko/sneklog/v2" ) @@ -19,6 +20,8 @@ type queuedObserverEvent struct { type observerDispatcher struct { observer Observer logger *sneklog.Logger + ctx context.Context + cancel context.CancelFunc queue chan queuedObserverEvent stop chan struct{} mu sync.RWMutex @@ -29,9 +32,12 @@ type observerDispatcher struct { } func newObserverDispatcher(observer Observer, logger *sneklog.Logger) *observerDispatcher { + ctx, cancel := context.WithCancel(context.Background()) dispatcher := &observerDispatcher{ observer: observer, logger: logger, + ctx: ctx, + cancel: cancel, queue: make(chan queuedObserverEvent, observerQueueSize), stop: make(chan struct{}), } @@ -43,9 +49,8 @@ func newObserverDispatcher(observer Observer, logger *sneklog.Logger) *observerD func (d *observerDispatcher) enqueue(ctx context.Context, event Event) { if ctx == nil { ctx = context.Background() - } else { - ctx = context.WithoutCancel(ctx) } + ctx = observerEventContext{Context: context.WithoutCancel(ctx), lifecycle: d.ctx} d.mu.RLock() defer d.mu.RUnlock() if d.closed { @@ -89,12 +94,32 @@ func (d *observerDispatcher) dispatch(queued queuedObserverEvent) { emitObserverEvent(d.observer, queued.ctx, queued.event) } -func (d *observerDispatcher) close() { +func (d *observerDispatcher) close(ctx context.Context) error { d.stopOnce.Do(func() { d.mu.Lock() d.closed = true + d.cancel() close(d.stop) d.mu.Unlock() }) - d.wg.Wait() + done := make(chan struct{}) + go func() { + d.wg.Wait() + close(done) + }() + select { + case <-done: + return nil + case <-ctx.Done(): + return fmt.Errorf("%w: %v", ErrObserverShutdownTimeout, ctx.Err()) + } } + +type observerEventContext struct { + context.Context + lifecycle context.Context +} + +func (ctx observerEventContext) Deadline() (time.Time, bool) { return ctx.lifecycle.Deadline() } +func (ctx observerEventContext) Done() <-chan struct{} { return ctx.lifecycle.Done() } +func (ctx observerEventContext) Err() error { return ctx.lifecycle.Err() } diff --git a/observer_async_test.go b/observer_async_test.go new file mode 100644 index 0000000..5b36197 --- /dev/null +++ b/observer_async_test.go @@ -0,0 +1,59 @@ +package laniakea + +import ( + "context" + "errors" + "testing" + "time" +) + +type cancelAwareObserver struct { + testObserver + started chan struct{} +} + +func (o *cancelAwareObserver) OnUpdateReceived(ctx context.Context, _ UpdateReceivedEvent) { + close(o.started) + <-ctx.Done() +} + +type stubbornObserver struct { + testObserver + started chan struct{} + release chan struct{} +} + +func (o *stubbornObserver) OnUpdateReceived(context.Context, UpdateReceivedEvent) { + close(o.started) + <-o.release +} + +func TestObserverDispatcherCancelsCallbackDuringClose(t *testing.T) { + observer := &cancelAwareObserver{started: make(chan struct{})} + dispatcher := newObserverDispatcher(observer, nil) + dispatcher.enqueue(context.Background(), UpdateReceivedEvent{}) + <-observer.started + + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + if err := dispatcher.close(ctx); err != nil { + t.Fatalf("close returned error: %v", err) + } +} + +func TestObserverDispatcherCloseTimeout(t *testing.T) { + observer := &stubbornObserver{started: make(chan struct{}), release: make(chan struct{})} + dispatcher := newObserverDispatcher(observer, nil) + dispatcher.enqueue(context.Background(), UpdateReceivedEvent{}) + <-observer.started + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) + defer cancel() + if err := dispatcher.close(ctx); !errors.Is(err, ErrObserverShutdownTimeout) { + t.Fatalf("close error = %v, want ErrObserverShutdownTimeout", err) + } + close(observer.release) + if err := dispatcher.close(context.Background()); err != nil { + t.Fatalf("second close returned error: %v", err) + } +} diff --git a/plugins.go b/plugins.go index 415c20c..76aca58 100644 --- a/plugins.go +++ b/plugins.go @@ -10,6 +10,8 @@ import ( "git.scuroneko.dev/scuroneko/sneklog/v2" ) +var errMiddlewareBlocked = errors.New("middleware blocked call") + // Plugin represents a collection of commands and payloads (e.g., callback handlers), // with shared middleware and configuration. // @@ -268,7 +270,7 @@ func (p *Plugin[T]) executeCmd(cmd string, ctx *MessageContext, db T) error { // Run command-specific middlewares for _, m := range command.middlewares { if !m.Execute(ctx, db) { - return AsInternalError(errors.New("middleware blocked call")) + return errMiddlewareBlocked } } @@ -289,7 +291,7 @@ func (p *Plugin[T]) executePayload(payload string, ctx *MessageContext, db T) er // Run command-specific middlewares for _, m := range command.middlewares { if !m.Execute(ctx, db) { - return AsInternalError(errors.New("middleware blocked call")) + return errMiddlewareBlocked } } @@ -363,6 +365,7 @@ func (m Middleware[T]) SetAsync(async bool) Middleware[T] { // continue to share storage with the synchronous flow. Async middleware // must treat those fields as read-only — mutating them races the sync chain // that mutates the same context concurrently. +// Bot runtimes wait for tracked asynchronous middleware before returning. func (m Middleware[T]) Execute(ctx *MessageContext, db T) bool { if m.executor == nil { reportMiddlewareError(ctx, m.name, ErrMiddlewareExecutorNil) @@ -370,14 +373,19 @@ func (m Middleware[T]) Execute(ctx *MessageContext, db T) bool { } if m.async { ctxCopy := *ctx - go func(ctx MessageContext) { + task := func() { defer func() { if recovered := recover(); recovered != nil { - reportMiddlewareError(&ctx, m.name, fmt.Errorf("%w in middleware %q: %v", ErrHandlerPanic, m.name, recovered)) + reportMiddlewareError(&ctxCopy, m.name, fmt.Errorf("%w in middleware %q: %v", ErrHandlerPanic, m.name, recovered)) } }() - m.executor(&ctx, db) - }(ctxCopy) + m.executor(&ctxCopy, db) + } + if ctx.asyncTask != nil { + ctx.asyncTask(task) + } else { + go task() + } return true } result := false diff --git a/plugins_test.go b/plugins_test.go index f003b98..b0e24d4 100644 --- a/plugins_test.go +++ b/plugins_test.go @@ -49,6 +49,39 @@ func TestAsyncMiddlewareRecoversPanic(t *testing.T) { } } +func TestAsyncMiddlewareUsesContextTaskTracker(t *testing.T) { + bot := new(Bot[NoData]) + started := make(chan struct{}) + release := make(chan struct{}) + ctx := &MessageContext{asyncTask: bot.startAsyncTask} + middleware := NewMiddleware[NoData]("tracked", func(*MessageContext, NoData) bool { + close(started) + <-release + return true + }).SetAsync(true) + + if !middleware.Execute(ctx, NoData{}) { + t.Fatal("async middleware blocked execution") + } + <-started + waited := make(chan struct{}) + go func() { + bot.middlewareWG.Wait() + close(waited) + }() + select { + case <-waited: + t.Fatal("task tracker finished before middleware returned") + case <-time.After(20 * time.Millisecond): + } + close(release) + select { + case <-waited: + case <-time.After(time.Second): + t.Fatal("task tracker did not finish after middleware returned") + } +} + func TestSyncMiddlewareRecoversPanic(t *testing.T) { observer := &middlewareErrorObserver{errors: make(chan ErrorEvent, 1)} ctx := &MessageContext{observer: observer} diff --git a/policy.go b/policy.go index 722c179..995af08 100644 --- a/policy.go +++ b/policy.go @@ -143,7 +143,7 @@ func RequireChatAdmin[T AppData]() Policy[T] { return AsInternalError(errors.New("chat-admin policy requires message chat context")) } - member, err := ctx.API.GetChatMember(tgapi.GetChatMember{ + member, err := ctx.API.GetChatMemberWithContext(ctx.Context(), tgapi.GetChatMember{ ChatID: ctx.ChatID, UserID: ctx.FromID, }) @@ -166,7 +166,7 @@ func RequireChatCreator[T AppData]() Policy[T] { return AsInternalError(errors.New("chat-creator policy requires message chat context")) } - member, err := ctx.API.GetChatMember(tgapi.GetChatMember{ + member, err := ctx.API.GetChatMemberWithContext(ctx.Context(), tgapi.GetChatMember{ ChatID: ctx.ChatID, UserID: ctx.FromID, }) @@ -192,7 +192,7 @@ func RequireBotAdmin[T AppData]() Policy[T] { return AsInternalError(errors.New("bot ID is not set in context")) } - member, err := ctx.API.GetChatMember(tgapi.GetChatMember{ + member, err := ctx.API.GetChatMemberWithContext(ctx.Context(), tgapi.GetChatMember{ ChatID: ctx.ChatID, UserID: ctx.botID, }) if err != nil { diff --git a/policy_test.go b/policy_test.go index 9338884..49dc750 100644 --- a/policy_test.go +++ b/policy_test.go @@ -157,6 +157,43 @@ func TestRequireChatAdminUsesNormalizedIDs(t *testing.T) { } } +func TestRequireChatAdminUsesMessageContextCancellation(t *testing.T) { + api := tgapi.NewAPI( + tgapi.NewAPIOpts("token"). + SetAPIURL("https://example.test"). + SetHTTPClient(&http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) { + if err := req.Context().Err(); err != nil { + return nil, err + } + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"application/json"}}, + Body: io.NopCloser(strings.NewReader(`{"ok":true,"result":{"status":"administrator","user":{"id":55,"is_bot":false,"first_name":"tester"}}}`)), + }, nil + })}), + ) + t.Cleanup(func() { + if err := api.Close(); err != nil { + t.Fatalf("Close returned error: %v", err) + } + }) + + requestCtx, cancel := context.WithCancel(context.Background()) + cancel() + ctx := &MessageContext{ + API: api, + ChatID: -2001, + FromID: 55, + Logger: sneklog.NewLogger(), + ctx: requestCtx, + } + + err := RequireChatAdmin[NoData]()(ctx, NoData{}) + if !errors.Is(err, context.Canceled) || !IsInternalError(err) { + t.Fatalf("RequireChatAdmin error = %v, want internal context.Canceled", err) + } +} + func TestAllPoliciesReturnsFirstError(t *testing.T) { want := AsUserError(errors.New("blocked")) policy := AllPolicies( diff --git a/runners.go b/runners.go index fa523fb..96d8e43 100644 --- a/runners.go +++ b/runners.go @@ -93,7 +93,7 @@ func (r Runner[T]) Every(timeout time.Duration) Runner[T] { // ExecRunners executes all runners registered on the Bot with context-based lifecycle management. // // Execution semantics by configuration: -// - every=0, async=true: Runs once in a goroutine (fire and forget). +// - every=0, async=true: Runs once in a goroutine; runtime shutdown waits for it. // - every=0, async=false: Runs once synchronously; warns if slower than 2 seconds. // - every>0, async=true: Runs in a loop with the configured interval until ctx.Done(). // - every>0, async=false: Skipped with a warning (invalid configuration). @@ -111,7 +111,7 @@ func (bot *Bot[T]) ExecRunners(ctx context.Context) { } if runner.every == 0 && runner.async { - // One-time async: fire and forget + // One-time async: non-blocking startup; runtime shutdown waits for completion. bot.runnerOnceWG.Add(1) go func(r Runner[T]) { defer bot.runnerOnceWG.Done() diff --git a/scene_handler.go b/scene_handler.go index 221c5ca..c31f4a1 100644 --- a/scene_handler.go +++ b/scene_handler.go @@ -1,24 +1,37 @@ package laniakea import ( - "errors" "fmt" "strings" "time" ) func (bot *Bot[T]) tryHandleScene(ctx *MessageContext) (bool, error) { - key, session, err := bot.findSceneSession(ctx) - if err != nil { - if errors.Is(err, ErrCantFindSession) { - return false, nil + for _, scope := range bot.sceneScopePriority { + key, ok := buildSceneKey(scope, ctx) + if !ok { + continue } - return false, err - } - if session.Scene == "" { - return false, nil - } + unlock := bot.sceneLocks.lock([]string{key}) + session, err := bot.sessionStore.Get(key) + if err != nil { + unlock() + return false, err + } + if session.Scene == "" { + unlock() + continue + } + + handled, err := bot.tryHandleSceneSession(ctx, key, session) + unlock() + return handled, err + } + return false, nil +} + +func (bot *Bot[T]) tryHandleSceneSession(ctx *MessageContext, key string, session SceneSession) (bool, error) { for _, plugin := range bot.plugins { scene, ok := plugin.scenes[session.Scene] if !ok { diff --git a/scene_locks.go b/scene_locks.go index 5b01f0d..947a776 100644 --- a/scene_locks.go +++ b/scene_locks.go @@ -68,13 +68,3 @@ func uniqueSortedStrings(values []string) []string { } return result } - -func sceneKeysForContext(ctx *MessageContext) []string { - keys := make([]string, 0, 3) - for _, scope := range []SceneScope{SceneScopeUserChat, SceneScopeChat, SceneScopeUser} { - if key, ok := buildSceneKey(scope, ctx); ok { - keys = append(keys, key) - } - } - return keys -} diff --git a/tgapi/api.go b/tgapi/api.go index c755c91..97f3a83 100644 --- a/tgapi/api.go +++ b/tgapi/api.go @@ -272,7 +272,7 @@ func (r TelegramRequest[R, P]) doRequest(ctx context.Context, api *API) (R, erro api.logger.Debugln("REQ", url, redactRequestLog(reqData)) resp, err := api.client.Do(req) if err != nil { - return zero, fmt.Errorf("HTTP request failed: %w", err) + return zero, fmt.Errorf("HTTP request failed: %w", redactHTTPError(err, api.token)) } respData, err := readBody(resp.Body) diff --git a/tgapi/log_redaction.go b/tgapi/log_redaction.go index a7f4a07..55ee03c 100644 --- a/tgapi/log_redaction.go +++ b/tgapi/log_redaction.go @@ -2,7 +2,9 @@ package tgapi import ( "encoding/json" + "errors" "fmt" + "net/url" "strings" ) @@ -55,3 +57,30 @@ func redactLogValue(value any) { func responseLogSummary(method string, size int) string { return fmt.Sprintf("method=%s bytes=%d body=omitted", method, size) } + +type redactedError struct { + err error + secret string +} + +func (e *redactedError) Error() string { + return strings.ReplaceAll(e.err.Error(), e.secret, redactedLogValue) +} + +func (e *redactedError) Unwrap() error { return e.err } + +func redactHTTPError(err error, token string) error { + if err == nil || token == "" || !strings.Contains(err.Error(), token) { + return err + } + + var urlErr *url.Error + if !errors.As(err, &urlErr) { + return &redactedError{err: err, secret: token} + } + + redactedURL := *urlErr + redactedURL.URL = strings.ReplaceAll(redactedURL.URL, token, redactedLogValue) + redactedURL.Err = &redactedError{err: urlErr.Err, secret: token} + return &redactedURL +} diff --git a/tgapi/log_redaction_test.go b/tgapi/log_redaction_test.go index 3d6a945..91182bd 100644 --- a/tgapi/log_redaction_test.go +++ b/tgapi/log_redaction_test.go @@ -1,6 +1,9 @@ package tgapi import ( + "errors" + "fmt" + "net/url" "strings" "testing" ) @@ -37,3 +40,29 @@ func TestResponseLogSummaryNeverContainsBody(t *testing.T) { t.Fatalf("response summary does not explain omission: %s", got) } } + +func TestRedactHTTPErrorRemovesTokenAndPreservesCause(t *testing.T) { + const token = "123456:secret-token" + cause := errors.New("transport failed") + original := &url.Error{ + Op: "Post", + URL: "https://api.telegram.org/bot" + token + "/sendMessage", + Err: fmt.Errorf("request for %s failed: %w", token, cause), + } + + got := redactHTTPError(original, token) + if strings.Contains(got.Error(), token) { + t.Fatalf("redacted HTTP error contains bot token: %v", got) + } + if !errors.Is(got, cause) { + t.Fatalf("redacted HTTP error lost its cause: %v", got) + } + + var gotURLError *url.Error + if !errors.As(got, &gotURLError) { + t.Fatalf("redacted HTTP error lost url.Error type: %T", got) + } + if strings.Contains(gotURLError.Error(), token) { + t.Fatalf("redacted url.Error contains bot token: %v", gotURLError) + } +} diff --git a/tgapi/methods.go b/tgapi/methods.go index f5eecf2..4b21a11 100644 --- a/tgapi/methods.go +++ b/tgapi/methods.go @@ -336,7 +336,7 @@ func (api *API) openFileByLink(ctx context.Context, link string) (io.ReadCloser, res, err := api.client.Do(req) if err != nil { - return nil, err + return nil, redactHTTPError(err, api.token) } if res.StatusCode < http.StatusOK || res.StatusCode >= http.StatusMultipleChoices { defer func() { diff --git a/tgapi/uploader_api.go b/tgapi/uploader_api.go index 3a94924..f90bc0b 100644 --- a/tgapi/uploader_api.go +++ b/tgapi/uploader_api.go @@ -156,7 +156,7 @@ func (r UploaderRequest[R, P]) doRequest(ctx context.Context, up *Uploader) (R, resp, err := up.api.client.Do(req) _ = requestBody.Close() if err != nil { - return zero, fmt.Errorf("HTTP upload request failed: %w", err) + return zero, fmt.Errorf("HTTP upload request failed: %w", redactHTTPError(err, up.api.token)) } body, err := readBody(resp.Body) diff --git a/tgrich/build_test.go b/tgrich/build_test.go index fc02318..db823a2 100644 --- a/tgrich/build_test.go +++ b/tgrich/build_test.go @@ -216,6 +216,13 @@ func TestBuildHTMLLimits(t *testing.T) { {"combined nesting too deep", func() []tgapi.InputRichBlock { return []tgapi.InputRichBlock{P(wrapBold(Text("text"), maxRichDepth))} }, ErrRichNestingTooDeep}, + {"rich text arrays too deep", func() []tgapi.InputRichBlock { + text := tgapi.RichText(Text("text")) + for range maxRichDepth { + text = tgapi.RichTextArray{text} + } + return []tgapi.InputRichBlock{P(text)} + }, ErrRichNestingTooDeep}, {"maximum block nesting", func() []tgapi.InputRichBlock { return []tgapi.InputRichBlock{wrapDetails(P(Text("text")), maxRichDepth-1)} }, nil}, diff --git a/tgrich/html_render.go b/tgrich/html_render.go index 2d81bba..ead7a55 100644 --- a/tgrich/html_render.go +++ b/tgrich/html_render.go @@ -80,9 +80,12 @@ func renderText(t tgapi.RichText, step int) (string, error) { case tgapi.RichTextPlain: return escapeHTML(string(el)), nil case tgapi.RichTextArray: + if step >= maxRichDepth { + return "", ErrRichNestingTooDeep + } var b strings.Builder for _, item := range el { - part, err := renderText(item, step) + part, err := renderText(item, step+1) if err != nil { return "", err } diff --git a/tgrich/validate.go b/tgrich/validate.go index e604160..46c16d6 100644 --- a/tgrich/validate.go +++ b/tgrich/validate.go @@ -119,8 +119,11 @@ func (v *richValidator) text(text tgapi.RichText, depth int) error { case tgapi.RichTextPlain: return v.addChars(string(el)) case tgapi.RichTextArray: + if depth >= maxRichDepth { + return ErrRichNestingTooDeep + } for _, item := range el { - if err := v.text(item, depth); err != nil { + if err := v.text(item, depth+1); err != nil { return err } } @@ -176,7 +179,7 @@ func (v *richValidator) text(text tgapi.RichText, depth int) error { } func validateAutomaticEntity(text tgapi.RichText, semantic string) error { - visible, err := visibleRichText(text) + visible, err := visibleRichText(text, 0) if err != nil { return err } @@ -186,16 +189,22 @@ func validateAutomaticEntity(text tgapi.RichText, semantic string) error { return nil } -func visibleRichText(text tgapi.RichText) (string, error) { +func visibleRichText(text tgapi.RichText, depth int) (string, error) { + if depth > maxRichDepth { + return "", ErrRichNestingTooDeep + } switch el := text.(type) { case nil: return "", nil case tgapi.RichTextPlain: return string(el), nil case tgapi.RichTextArray: + if depth >= maxRichDepth { + return "", ErrRichNestingTooDeep + } var result strings.Builder for _, item := range el { - part, err := visibleRichText(item) + part, err := visibleRichText(item, depth+1) if err != nil { return "", err } @@ -209,33 +218,33 @@ func visibleRichText(text tgapi.RichText) (string, error) { case tgapi.RichTextAnchor: return "", nil case tgapi.RichTextWrap: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextURL: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextEmailAddress: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextPhoneNumber: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextBankCardNumber: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextMention: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextHashtag: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextCashtag: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextBotCommand: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextAnchorLink: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextReference: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextReferenceLink: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextDateTime: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) case tgapi.RichTextTextMention: - return visibleRichText(el.Text) + return visibleRichText(el.Text, depth+1) default: return "", ErrRichUnknownTag } diff --git a/utils/limiter.go b/utils/limiter.go index 659c5c2..490c6e5 100644 --- a/utils/limiter.go +++ b/utils/limiter.go @@ -28,7 +28,8 @@ type RateLimiter struct { chatLocks map[int64]time.Time // per-chat cooldown timestamps chatLimiters map[int64]*rate.Limiter // per-chat token buckets (1 req/sec) chatLastSeen map[int64]time.Time // last access timestamp per chat, for Cleanup eviction - chatMu sync.RWMutex // protects chatLocks, chatLimiters, and chatLastSeen + chatActive map[int64]int // in-flight users of each per-chat limiter + chatMu sync.RWMutex // protects all per-chat maps } // NewRateLimiter creates a new RateLimiter with default limits. @@ -40,6 +41,7 @@ func NewRateLimiter() *RateLimiter { chatLimiters: make(map[int64]*rate.Limiter), chatLocks: make(map[int64]time.Time), chatLastSeen: make(map[int64]time.Time), + chatActive: make(map[int64]int), } } @@ -58,11 +60,12 @@ func (rl *RateLimiter) Cleanup(idleThreshold time.Duration) { defer rl.chatMu.Unlock() for chatID, lastSeen := range rl.chatLastSeen { - if now.Sub(lastSeen) <= idleThreshold { + if now.Sub(lastSeen) <= idleThreshold || rl.chatActive[chatID] > 0 { continue } delete(rl.chatLimiters, chatID) delete(rl.chatLastSeen, chatID) + delete(rl.chatActive, chatID) } for chatID, until := range rl.chatLocks { if !until.After(now) { @@ -131,7 +134,8 @@ func (rl *RateLimiter) Wait(ctx context.Context, chatID int64) error { if err := rl.waitForChatUnlock(ctx, chatID); err != nil { return err } - chatLimiter := rl.getChatLimiter(chatID) + chatLimiter, release := rl.acquireChatLimiter(chatID) + defer release() if err := chatLimiter.Wait(ctx); err != nil { return err } @@ -200,7 +204,8 @@ func (rl *RateLimiter) Allow(chatID int64) bool { } } - chatLimiter := rl.getChatLimiter(chatID) + chatLimiter, release := rl.acquireChatLimiter(chatID) + defer release() chatReservation := chatLimiter.ReserveN(now, 1) if !chatReservation.OK() || chatReservation.DelayFrom(now) > 0 { chatReservation.CancelAt(now) @@ -311,3 +316,22 @@ func (rl *RateLimiter) getChatLimiter(chatID int64) *rate.Limiter { rl.chatLastSeen[chatID] = now return lim } + +func (rl *RateLimiter) acquireChatLimiter(chatID int64) (*rate.Limiter, func()) { + rl.chatMu.Lock() + limiter, ok := rl.chatLimiters[chatID] + if !ok { + limiter = rate.NewLimiter(1, 1) + rl.chatLimiters[chatID] = limiter + } + rl.chatLastSeen[chatID] = time.Now() + rl.chatActive[chatID]++ + rl.chatMu.Unlock() + + return limiter, func() { + rl.chatMu.Lock() + rl.chatActive[chatID]-- + rl.chatLastSeen[chatID] = time.Now() + rl.chatMu.Unlock() + } +} diff --git a/utils/limiter_test.go b/utils/limiter_test.go index 88e837d..a7ae728 100644 --- a/utils/limiter_test.go +++ b/utils/limiter_test.go @@ -189,3 +189,18 @@ func TestRateLimiterCleanupEvictsIdleChats(t *testing.T) { t.Fatal("expected future chat 11 lock to remain") } } + +func TestRateLimiterCleanupPreservesAcquiredLimiter(t *testing.T) { + rl := NewRateLimiter() + limiter, release := rl.acquireChatLimiter(42) + + rl.chatMu.Lock() + rl.chatLastSeen[42] = time.Now().Add(-time.Hour) + rl.chatMu.Unlock() + rl.Cleanup(time.Minute) + + if got := rl.getChatLimiter(42); got != limiter { + t.Fatal("Cleanup replaced a limiter while it was acquired") + } + release() +}