(new): PollTimeout config, RateLimiter.Cleanup
Golang lint / lint (pull_request) Successful in 3m4s
Golang lint / lint (push) Successful in 3m5s

(fix): compact payload escape, Draft.push validation, plugin logger ownership, worker StopAndWait, getChatLimiter deadlock, runner ctx-after-tick
(refactor): remove NewPayload, buildSceneKey from sceneRuntime, unify ToJSON fallback
(tests): compact round-trip, Draft.push state, RateLimiter.Cleanup eviction
(doc): changelog v1.0.0 rewrite, NewCommand dual-use godoc
This commit is contained in:
2026-05-18 17:27:14 +03:00
parent 09fb9261df
commit affb802a7b
38 changed files with 697 additions and 345 deletions
+2
View File
@@ -4,3 +4,5 @@
test/ test/
.codex/ .codex/
.codex .codex
.agents/
.claude/
+1 -1
View File
@@ -1,7 +1,7 @@
# AGENTS.md # AGENTS.md
## Purpose ## Purpose
This repository uses Codex for full-project Go code review, not diff-only review. This repository uses AI coding agents for full-project Go code review, not diff-only review.
When asked to review code, inspect the entire repository and use repository-wide context. Do not limit analysis to the latest commit, pull request diff, or recently changed files. When asked to review code, inspect the entire repository and use repository-wide context. Do not limit analysis to the latest commit, pull request diff, or recently changed files.
+27 -4
View File
@@ -3,28 +3,47 @@
## v1.0.0 ## v1.0.0
### Breaking Changes ### Breaking Changes
- Renamed `MsgContext` to `MessageContext` across the public API, including handler signatures (`CommandExecutor`, `MiddlewareExecutor`, scene handler types), all reply/edit/scene helpers, embedded fields on `SceneContext`, and documentation.
- Removed the `NewPayload(...)` constructor. `NewCommand(...)` builds the underlying `Command[T]` for both `/-`commands and callback payloads; registration via `Plugin.AddPayload`/`Plugin.Payload` decides routing.
- `MessageContext.Error(...)` no longer sends unclassified errors to the user. Only errors marked with `AsUserError(...)` are surfaced through the centralized reply path; everything else stays internal-only and is logged.
- `Plugin.Close()` no longer closes a logger supplied through `Plugin.SetLogger(...)`. Only loggers created by the bot during `AddPlugins` registration are owned and closed; caller-supplied loggers remain the caller's responsibility.
- Renamed final public APIs to idiomatic names before the stable release: `RunWebhookWithContext(...)`, `RunWebhook(...)`, `CloseWebhook()`, `BotWebhookOpts`, `NewBotWebhookOpts()`, `SetWebhookLogger(...)`, and `GetWebhookLogger()`. - Renamed final public APIs to idiomatic names before the stable release: `RunWebhookWithContext(...)`, `RunWebhook(...)`, `CloseWebhook()`, `BotWebhookOpts`, `NewBotWebhookOpts()`, `SetWebhookLogger(...)`, and `GetWebhookLogger()`.
- Renamed plugin builder helpers from `NewCommand(...)`, `NewPayload(...)`, and `NewScene(...)` to `Command(...)`, `Payload(...)`, and `Scene(...)`; `NewCommand(...)` and `NewPayload(...)` now take the command string before the executor. - Renamed plugin builder helpers from `NewCommand(...)` and `NewScene(...)` to `Command(...)` and `Scene(...)`; the surviving `NewCommand(...)` takes the command string before the executor.
- Renamed command argument value constants to `CommandValueString`, `CommandValueInt`, `CommandValueBool`, and `CommandValueAny`; `NewCommandArg(...)` now defaults to unvalidated `CommandValueAny`. - Renamed command argument value constants to `CommandValueString`, `CommandValueInt`, `CommandValueBool`, and `CommandValueAny`; `NewCommandArg(...)` now defaults to unvalidated `CommandValueAny`.
- Renamed runner builders from `Onetime(...)` and `Timeout(...)` to `Once(...)` and `Every(...)`. - Renamed runner builders from `Onetime(...)` and `Timeout(...)` to `Once(...)` and `Every(...)`.
- Renamed remaining public acronym/casing outliers including `AnswerCallback...`, `ParseMarkdownV2`, `ParseMarkdown`, `GetChatMemberCount`, `DropRateLimitOverflow`, `SetDropRateLimitOverflow`, and inline keyboard builder APIs. - Renamed remaining public acronym/casing outliers including `AnswerCallback...`, `ParseMarkdownV2`, `ParseMarkdown`, `GetChatMemberCount`, `DropRateLimitOverflow`, `SetDropRateLimitOverflow`, and inline keyboard builder APIs.
### Added ### Added
- Added `MsgContext.IsCallback()` and `MsgContext.HasPhoto()` helpers for callback-aware handler code. - Added `MessageContext.IsCallback()` and `MessageContext.HasPhoto()` helpers for callback-aware handler code.
- Added `MsgContext.UpsertKeyboard(...)` and `MsgContext.UpsertKeyboardMarkdown(...)` helpers that edit callback messages, replace photo callback messages with a fresh chat message, and send a new chat message outside callback flow. - Added `MessageContext.UpsertKeyboard(...)` and `MessageContext.UpsertKeyboardMarkdown(...)` helpers that edit callback messages, replace photo callback messages with a fresh chat message, and send a new chat message outside callback flow.
- Added `CommandGroup`, `NewCommandGroup(...)`, `Plugin.CommandGroup(...)`, and `Plugin.AddCommandGroup(...)` helpers for registering prefixed command groups with shared middleware. - Added `CommandGroup`, `NewCommandGroup(...)`, `Plugin.CommandGroup(...)`, and `Plugin.AddCommandGroup(...)` helpers for registering prefixed command groups with shared middleware.
- Added the `tgfmt` package with typed MarkdownV2, HTML, legacy Markdown formatting helpers, and a message entity builder. - Added the `tgfmt` package with typed MarkdownV2, HTML, legacy Markdown formatting helpers, and a message entity builder.
- Added `InlineKeyboardButtonBuilder.SetPayloadType(...)`, `InlineKeyboardButtonBuilder.SetCallbackData(...)`, and `MsgContext.NewInlineKeyboardButton(...)` helpers for payload-aware button building. - Added `InlineKeyboardButtonBuilder.SetPayloadType(...)`, `InlineKeyboardButtonBuilder.SetCallbackData(...)`, and `MessageContext.NewInlineKeyboardButton(...)` helpers for payload-aware button building.
- Added compact callback payload encoding through `BotPayloadCompact`, `BotPayloadCompactBase64`, compact inline keyboard builders, and matching `CallbackData` helpers. - Added compact callback payload encoding through `BotPayloadCompact`, `BotPayloadCompactBase64`, compact inline keyboard builders, and matching `CallbackData` helpers.
- Added `BotOpts.PollTimeout`, `BotOpts.SetPollTimeout(...)`, and the `POLL_TIMEOUT` environment variable to configure the long-polling `getUpdates` timeout (default 30 seconds).
- Added `RateLimiter.Cleanup(idleThreshold)` to evict per-chat limiter state and expired chat cooldowns; the limiter now tracks per-chat last-seen time so long-running bots can bound memory through a periodic runner.
- Added cached bot identity (`Bot.userID`) populated at `NewBot` so chat-admin policies and similar lookups reuse it instead of issuing a fresh `GetMe` request.
- Added `tgapi.ResponseError` so Telegram API error codes, descriptions, and response parameters remain inspectable through returned errors. - Added `tgapi.ResponseError` so Telegram API error codes, descriptions, and response parameters remain inspectable through returned errors.
### Changed ### Changed
- Version metadata now reports the stable `v1.0.0` release instead of `v1.0.0-rc.16`. - Version metadata now reports the stable `v1.0.0` release instead of `v1.0.0-rc.16`.
- Compact callback payload encoding now escapes `,`, `|`, and `\` in command and arg bytes so payloads containing those bytes round-trip without ambiguity. Note: the format coalesces "no args" with "single empty arg" — both encode as `cmd|` and decode to nil args.
- `CallbackData.ToJSON()`, `ToBase64()`, `ToCompact()`, and `ToCompactBase64()` now all return an empty string on serialization failure; the previous `ToJSON()` fallback `{"cmd":""}` has been removed so encoder bugs surface visibly instead of routing to no handler.
- Bot-level middleware blocks now emit a final `UpdateHandledEvent` with `Handled=false`, keeping observer update lifecycles balanced. - Bot-level middleware blocks now emit a final `UpdateHandledEvent` with `Handled=false`, keeping observer update lifecycles balanced.
- Plugin registration now warns when `AddCommand`, `AddPayload`, or `AddScene` overwrites an existing entry with the same name instead of silently replacing it.
- `BotOpts`, `tgapi.APIOpts`, logger utilities, README, and wiki pages now document the final stable API names and configuration options consistently. - `BotOpts`, `tgapi.APIOpts`, logger utilities, README, and wiki pages now document the final stable API names and configuration options consistently.
- CI now checks formatting, tests, vet, and lint on both pushes and pull requests. - CI now checks formatting, tests, vet, and lint on both pushes and pull requests.
### Fixed ### Fixed
- Fixed the update worker pool returning before in-flight handlers completed. `startUpdateWorkers` now calls `pool.StopAndWait()` so the bot waits for already-submitted tasks before runtime exit.
- Fixed `RateLimiter.getChatLimiter` upgrading a held read lock to a write lock, which could deadlock under contention. The lookup now releases the read lock before acquiring the write lock and re-checks the map.
- Fixed `RateLimiter` per-chat limiter and lock maps growing unbounded for the lifetime of long-running bots that serve many distinct chats.
- Fixed `Draft.Push` mutating `Message` before validating the candidate length, leaving the draft in a half-mutated state when the candidate would exceed Telegram's limit. The candidate is now validated first; on failure the draft remains unchanged.
- Fixed background runners running one extra iteration after context cancellation when both `ctx.Done()` and the ticker were ready in the same `select`.
- Fixed `Plugin.Close()` double-closing a logger supplied by the caller through `SetLogger(...)`.
- Fixed compact callback payload corruption for arguments containing `,` or `|` bytes.
- Fixed `LoadOptsFromEnv` calling `os.Getenv("MAX_WORKERS")` twice when parsing the worker count.
- Fixed `sceneRuntime` interface carrying a delegating `buildSceneKey` method that just forwarded to a package-level helper; `MessageContext` scene helpers now call the helper directly.
- Fixed webhook startup so empty-secret warnings are logged only after the webhook logger is initialized. - Fixed webhook startup so empty-secret warnings are logged only after the webhook logger is initialized.
- Fixed webhook startup so a logger configured through `SetWebhookLogger(...)` is preserved. - Fixed webhook startup so a logger configured through `SetWebhookLogger(...)` is preserved.
- Fixed long-polling 429 handling so `getUpdates` retries use Telegram `retry_after` directly and do not inflate later transient-error backoff. - Fixed long-polling 429 handling so `getUpdates` retries use Telegram `retry_after` directly and do not inflate later transient-error backoff.
@@ -38,6 +57,10 @@
- Added regression coverage for context-aware inline keyboard button payload encoding. - Added regression coverage for context-aware inline keyboard button payload encoding.
- Added regression coverage for compact and Base64-encoded compact callback payload decoding. - Added regression coverage for compact and Base64-encoded compact callback payload decoding.
- Added regression coverage for long-polling `retry_after` handling on Telegram 429 responses. - Added regression coverage for long-polling `retry_after` handling on Telegram 429 responses.
- Added regression coverage for compact callback payload round-tripping through `,`, `|`, and `\` separator bytes and a missing-separator decode error.
- Added regression coverage for `Draft.Push` preserving the existing message when validation rejects the candidate.
- Added regression coverage for `RateLimiter.Cleanup` evicting idle chat limiters and expired chat locks while leaving active state in place.
- Updated `MessageContext.Error` tests so unclassified errors stay internal-only and only `AsUserError` reaches the user.
## v1.0.0-rc.16 ## v1.0.0-rc.16
+10 -1
View File
@@ -36,7 +36,7 @@ type AppData any
// data. // data.
// //
// Use Bot[NoData] to indicate no shared dependency injection is required. // Use Bot[NoData] to indicate no shared dependency injection is required.
type NoData struct{ AppData } type NoData struct{}
// AppDataLogger builds a sneklog.LoggerWriter from injected application data. // AppDataLogger builds a sneklog.LoggerWriter from injected application data.
// //
@@ -89,10 +89,12 @@ type Bot[T AppData] struct {
token string token string
debug bool debug bool
errorTemplate string errorTemplate string
userID int64
username string username string
payloadType BotPayloadType payloadType BotPayloadType
strictPayloadType bool strictPayloadType bool
maxWorkers int maxWorkers int
pollTimeout int // Long-polling timeout in seconds for getUpdates
logFormat utils.LogFormat logFormat utils.LogFormat
logFormatter *sneklog.Formatter logFormatter *sneklog.Formatter
@@ -184,12 +186,18 @@ func NewBot[T any](opts *BotOpts) (*Bot[T], error) {
workers = opts.MaxWorkers workers = opts.MaxWorkers
} }
pollTimeout := 30
if opts.PollTimeout > 0 {
pollTimeout = opts.PollTimeout
}
bot := &Bot[T]{ bot := &Bot[T]{
updateOffset: 0, updateOffset: 0,
errorTemplate: "%s", errorTemplate: "%s",
payloadType: BotPayloadBase64, payloadType: BotPayloadBase64,
strictPayloadType: opts.StrictPayloadType, strictPayloadType: opts.StrictPayloadType,
maxWorkers: workers, maxWorkers: workers,
pollTimeout: pollTimeout,
updateQueue: updateQueue, updateQueue: updateQueue,
api: api, api: api,
uploader: uploader, uploader: uploader,
@@ -239,6 +247,7 @@ func NewBot[T any](opts *BotOpts) (*Bot[T], error) {
return nil, err return nil, err
} }
bot.username = Val(u.Username, "") bot.username = Val(u.Username, "")
bot.userID = u.ID
if bot.username == "" { if bot.username == "" {
bot.logger.Warn("Can't get bot username. Named command handlers won't work!") bot.logger.Warn("Can't get bot username. Named command handlers won't work!")
} }
+2 -2
View File
@@ -60,7 +60,7 @@ func (bot *Bot[T]) SetSessionStore(store SessionStore) *Bot[T] {
return bot return bot
} }
if store == nil { if store == nil {
bot.logger.Warn("SetSessionStore called with nil store; using default MemorySessionStore") bot.logger.Warn("SetSessionStore called with nil store; nothing changed")
return bot return bot
} }
bot.sessionStore = store bot.sessionStore = store
@@ -216,7 +216,7 @@ func (bot *Bot[T]) SetL10n(l *L10n) *Bot[T] {
return bot return bot
} }
if l == nil { if l == nil {
bot.logger.Warn("SetL10n called with nil L10n; localization will be disabled") bot.logger.Warn("SetL10n called with nil L10n; localization will not change")
return bot return bot
} }
bot.l10n = l bot.l10n = l
+22 -1
View File
@@ -65,6 +65,11 @@ type BotOpts struct {
// MaxWorkers is the maximum number of update handlers that may run concurrently. // MaxWorkers is the maximum number of update handlers that may run concurrently.
MaxWorkers int MaxWorkers int
// PollTimeout is the long-polling timeout in seconds for getUpdates.
// Defaults to 30. Telegram allows 0..50; values outside that range are accepted
// by the bot but rejected by Telegram at runtime.
PollTimeout int
// FileConfigVersion stores the version declared by the config file used to // FileConfigVersion stores the version declared by the config file used to
// load these options. // load these options.
// //
@@ -94,6 +99,7 @@ type BotOpts struct {
// - DROP_RL_OVERFLOW: "true" to drop updates on rate limit overflow // - DROP_RL_OVERFLOW: "true" to drop updates on rate limit overflow
// - STRICT_PAYLOAD_TYPE: "true" to reject callback payloads encoded in a different format // - STRICT_PAYLOAD_TYPE: "true" to reject callback payloads encoded in a different format
// - MAX_WORKERS: maximum number of concurrent update handlers (default: 32) // - MAX_WORKERS: maximum number of concurrent update handlers (default: 32)
// - POLL_TIMEOUT: long-polling timeout in seconds for getUpdates (default: 30)
// - LOG_FORMAT: logger output format, "text" or "json" (default: "text") // - LOG_FORMAT: logger output format, "text" or "json" (default: "text")
// //
// Returns a populated BotOpts. // Returns a populated BotOpts.
@@ -101,6 +107,7 @@ type BotOpts struct {
func LoadOptsFromEnv() *BotOpts { func LoadOptsFromEnv() *BotOpts {
rateLimit := 30 rateLimit := 30
maxWorkers := 32 maxWorkers := 32
pollTimeout := 30
stringUpdateTypes := splitEnvList(os.Getenv("UPDATE_TYPES")) stringUpdateTypes := splitEnvList(os.Getenv("UPDATE_TYPES"))
updateTypes := make([]tgapi.UpdateType, 0, len(stringUpdateTypes)) updateTypes := make([]tgapi.UpdateType, 0, len(stringUpdateTypes))
@@ -115,11 +122,17 @@ func LoadOptsFromEnv() *BotOpts {
} }
if mw := os.Getenv("MAX_WORKERS"); mw != "" { if mw := os.Getenv("MAX_WORKERS"); mw != "" {
if n, err := strconv.Atoi(os.Getenv("MAX_WORKERS")); err == nil { if n, err := strconv.Atoi(mw); err == nil {
maxWorkers = n maxWorkers = n
} }
} }
if pt := os.Getenv("POLL_TIMEOUT"); pt != "" {
if n, err := strconv.Atoi(pt); err == nil {
pollTimeout = n
}
}
return &BotOpts{ return &BotOpts{
Token: os.Getenv("TG_TOKEN"), Token: os.Getenv("TG_TOKEN"),
UpdateTypes: updateTypes, UpdateTypes: updateTypes,
@@ -140,6 +153,7 @@ func LoadOptsFromEnv() *BotOpts {
StrictPayloadType: os.Getenv("STRICT_PAYLOAD_TYPE") == "true", StrictPayloadType: os.Getenv("STRICT_PAYLOAD_TYPE") == "true",
MaxWorkers: maxWorkers, MaxWorkers: maxWorkers,
PollTimeout: pollTimeout,
FileConfigVersion: 0, FileConfigVersion: 0,
LogFormat: utils.LogFormat(os.Getenv("LOG_FORMAT")), LogFormat: utils.LogFormat(os.Getenv("LOG_FORMAT")),
} }
@@ -256,6 +270,13 @@ func (opts *BotOpts) SetMaxWorkers(workers int) *BotOpts {
return opts return opts
} }
// SetPollTimeout sets the long-polling timeout in seconds for getUpdates.
// Defaults to 30. Telegram accepts 0..50.
func (opts *BotOpts) SetPollTimeout(seconds int) *BotOpts {
opts.PollTimeout = seconds
return opts
}
// SetLogFormat sets the output format used by bot-managed loggers. // SetLogFormat sets the output format used by bot-managed loggers.
func (opts *BotOpts) SetLogFormat(format utils.LogFormat) *BotOpts { func (opts *BotOpts) SetLogFormat(format utils.LogFormat) *BotOpts {
opts.LogFormat = format opts.LogFormat = format
+1 -1
View File
@@ -155,7 +155,7 @@ func SaveBotOptsFile(codec BotOptsFileCodec, filename string, opts *BotOpts) err
if err != nil { if err != nil {
return err return err
} }
err = os.WriteFile(filename, data, 0644) err = os.WriteFile(filename, data, 0600)
if err != nil { if err != nil {
return err return err
} }
+1
View File
@@ -29,6 +29,7 @@ func (bot *Bot[T]) AddPlugins(plugin ...*Plugin[T]) *Bot[T] {
cloned := clonePlugin(p) cloned := clonePlugin(p)
if cloned.logger == nil { if cloned.logger == nil {
cloned.logger = utils.CreateLogger(cloned.name, level, bot.logFormat, bot.logFormatter) cloned.logger = utils.CreateLogger(cloned.name, level, bot.logFormat, bot.logFormatter)
cloned.loggerOwned = true
} }
bot.addTokenReplacer(cloned.logger) bot.addTokenReplacer(cloned.logger)
bot.plugins = append(bot.plugins, cloned) bot.plugins = append(bot.plugins, cloned)
+1 -5
View File
@@ -33,7 +33,7 @@ func (bot *Bot[T]) findScene(name string) (*sceneMeta, bool) {
return nil, false return nil, false
} }
func (bot *Bot[T]) findSceneSession(ctx *MsgContext) (string, SceneSession, error) { func (bot *Bot[T]) findSceneSession(ctx *MessageContext) (string, SceneSession, error) {
var zero SceneSession var zero SceneSession
for _, scope := range bot.sceneScopePriority { for _, scope := range bot.sceneScopePriority {
@@ -53,7 +53,3 @@ func (bot *Bot[T]) findSceneSession(ctx *MsgContext) (string, SceneSession, erro
return "", zero, ErrCantFindSession return "", zero, ErrCantFindSession
} }
func (bot *Bot[T]) buildSceneKey(scope SceneScope, ctx *MsgContext) (string, bool) {
return buildSceneKey(scope, ctx)
}
+6 -6
View File
@@ -63,14 +63,14 @@ func TestAddPluginsSnapshotsConfiguration(t *testing.T) {
bot := &Bot[NoData]{logger: sneklog.NewLogger()} bot := &Bot[NoData]{logger: sneklog.NewLogger()}
plugin := NewPlugin[NoData]("demo") plugin := NewPlugin[NoData]("demo")
cmd := plugin.Command("start", func(ctx *MsgContext, db NoData) error { return nil }) cmd := plugin.Command("start", func(ctx *MessageContext, db NoData) error { return nil })
plugin.AddMiddleware(NewMiddleware("base", func(ctx *MsgContext, db NoData) bool { return true })) plugin.AddMiddleware(NewMiddleware("base", func(ctx *MessageContext, db NoData) bool { return true }))
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
cmd.SetDescription("mutated after registration") cmd.SetDescription("mutated after registration")
plugin.Command("late", func(ctx *MsgContext, db NoData) error { return nil }) plugin.Command("late", func(ctx *MessageContext, db NoData) error { return nil })
plugin.AddMiddleware(NewMiddleware("late", func(ctx *MsgContext, db NoData) bool { return true })) plugin.AddMiddleware(NewMiddleware("late", func(ctx *MessageContext, db NoData) bool { return true }))
registered := bot.plugins[0] registered := bot.plugins[0]
if _, exists := registered.commands["late"]; exists { if _, exists := registered.commands["late"]; exists {
@@ -812,7 +812,7 @@ func TestAddPluginsAndRuntimeRegistrationsNoOpAfterRunStarts(t *testing.T) {
bot := &Bot[NoData]{ bot := &Bot[NoData]{
logger: sneklog.NewLogger(), logger: sneklog.NewLogger(),
prefixes: []string{"/"}, prefixes: []string{"/"},
middlewares: []Middleware[NoData]{NewMiddleware("base", func(ctx *MsgContext, db NoData) bool { return true })}, middlewares: []Middleware[NoData]{NewMiddleware("base", func(ctx *MessageContext, db NoData) bool { return true })},
runners: []Runner[NoData]{NewRunner("base", func(bot *Bot[NoData]) error { return nil })}, runners: []Runner[NoData]{NewRunner("base", func(bot *Bot[NoData]) error { return nil })},
} }
plugin := NewPlugin[NoData]("late") plugin := NewPlugin[NoData]("late")
@@ -823,7 +823,7 @@ func TestAddPluginsAndRuntimeRegistrationsNoOpAfterRunStarts(t *testing.T) {
defer bot.finishRun() defer bot.finishRun()
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
bot.AddMiddleware(NewMiddleware("late", func(ctx *MsgContext, db NoData) bool { return true })) bot.AddMiddleware(NewMiddleware("late", func(ctx *MessageContext, db NoData) bool { return true }))
bot.AddRunner(NewRunner("late", func(bot *Bot[NoData]) error { return nil })) bot.AddRunner(NewRunner("late", func(bot *Bot[NoData]) error { return nil }))
if len(bot.plugins) != 0 { if len(bot.plugins) != 0 {
+2 -1
View File
@@ -68,7 +68,7 @@ func (bot *Bot[T]) startUpdateWorkers(ctx context.Context) {
bot.handle(ctx, u) bot.handle(ctx, u)
}) })
} }
pool.Stop() // Wait for all tasks to complete and stop the pool pool.StopAndWait() // Wait for all tasks to complete and stop the pool
} }
func (bot *Bot[T]) initLoggers(opts *BotOpts) { func (bot *Bot[T]) initLoggers(opts *BotOpts) {
@@ -183,6 +183,7 @@ func clonePlugin[T AppData](p *Plugin[T]) Plugin[T] {
middlewares: append(extypes.Slice[Middleware[T]](nil), p.middlewares...), middlewares: append(extypes.Slice[Middleware[T]](nil), p.middlewares...),
skipAutoCmd: p.skipAutoCmd, skipAutoCmd: p.skipAutoCmd,
logger: p.logger, logger: p.logger,
loggerOwned: false, // user-supplied loggers stay caller-owned; bot may take ownership during registration
messageFallback: p.messageFallback, messageFallback: p.messageFallback,
handlers: make(map[tgapi.UpdateType]CommandExecutor[T]), handlers: make(map[tgapi.UpdateType]CommandExecutor[T]),
onClose: p.onClose, onClose: p.onClose,
+15 -34
View File
@@ -370,21 +370,14 @@ func (bot *Bot[T]) newWebhookMux(ctx context.Context, opts *BotWebhookOpts) *htt
r.HandleFunc(opts.Path, updateHandler(ctx, bot, opts.SecretToken)) r.HandleFunc(opts.Path, updateHandler(ctx, bot, opts.SecretToken))
return r return r
} }
func (bot *Bot[T]) runWebhook(ctx context.Context, opts *BotWebhookOpts) error { func (bot *Bot[T]) baseRunWebhook(ctx context.Context, opts *BotWebhookOpts, runFunc func(*http.Server, chan error)) error {
srv := &http.Server{ srv := &http.Server{
Addr: fmt.Sprintf(":%d", opts.LocalPort), Addr: fmt.Sprintf(":%d", opts.LocalPort),
Handler: bot.newWebhookMux(ctx, opts), Handler: bot.newWebhookMux(ctx, opts),
} }
errCh := make(chan error, 1) errCh := make(chan error, 1)
go func() { go runFunc(srv, errCh)
err := srv.ListenAndServe()
if err != nil && !errors.Is(err, http.ErrServerClosed) {
errCh <- err
return
}
errCh <- nil
}()
bot.webhookLogger.Infoln(fmt.Sprintf("Bot Webhook started at %s; waiting for updates at %s", srv.Addr, opts.URL)) bot.webhookLogger.Infoln(fmt.Sprintf("Bot Webhook started at %s; waiting for updates at %s", srv.Addr, opts.URL))
@@ -403,38 +396,26 @@ func (bot *Bot[T]) runWebhook(ctx context.Context, opts *BotWebhookOpts) error {
return err return err
} }
} }
func (bot *Bot[T]) runWebhookTLS(ctx context.Context, opts *BotWebhookOpts, key, cert string) error { func (bot *Bot[T]) runWebhook(ctx context.Context, opts *BotWebhookOpts) error {
srv := &http.Server{ return bot.baseRunWebhook(ctx, opts, func(srv *http.Server, errCh chan error) {
Addr: fmt.Sprintf(":%d", opts.LocalPort), err := srv.ListenAndServe()
Handler: bot.newWebhookMux(ctx, opts), if err != nil && !errors.Is(err, http.ErrServerClosed) {
} errCh <- err
errCh := make(chan error, 1) return
}
errCh <- nil
})
go func() { }
func (bot *Bot[T]) runWebhookTLS(ctx context.Context, opts *BotWebhookOpts, key, cert string) error {
return bot.baseRunWebhook(ctx, opts, func(srv *http.Server, errCh chan error) {
err := srv.ListenAndServeTLS(cert, key) err := srv.ListenAndServeTLS(cert, key)
if err != nil && !errors.Is(err, http.ErrServerClosed) { if err != nil && !errors.Is(err, http.ErrServerClosed) {
errCh <- err errCh <- err
return return
} }
errCh <- nil errCh <- nil
}() })
bot.webhookLogger.Infoln(fmt.Sprintf("Bot webhook started with TLS(%s, %s) at %s; waiting for updates at %s", key, cert, srv.Addr, opts.URL))
select {
case <-ctx.Done():
shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
return err
}
return <-errCh
case err := <-errCh:
return err
}
} }
func validateWebhookPath(path string, useStatusPath bool) error { func validateWebhookPath(path string, useStatusPath bool) error {
if path == "" { if path == "" {
+1 -1
View File
@@ -135,7 +135,7 @@ func TestRunWebhookRuntimePreservesConfiguredWebhookLogger(t *testing.T) {
func TestRunWebhookRuntimeProcessesEnqueuedUpdate(t *testing.T) { func TestRunWebhookRuntimeProcessesEnqueuedUpdate(t *testing.T) {
var calls atomic.Int32 var calls atomic.Int32
plugin := NewPlugin[NoData]("demo") plugin := NewPlugin[NoData]("demo")
plugin.Command("start", func(ctx *MsgContext, db NoData) error { plugin.Command("start", func(ctx *MessageContext, db NoData) error {
calls.Add(1) calls.Add(1)
return nil return nil
}) })
+2 -2
View File
@@ -44,7 +44,7 @@ func TestAutoGenerateCommandsChecksLimitBeforeDelete(t *testing.T) {
}() }()
plugin := NewPlugin[NoData]("overflow") plugin := NewPlugin[NoData]("overflow")
exec := func(ctx *MsgContext, db NoData) error { return nil } exec := func(ctx *MessageContext, db NoData) error { return nil }
for i := 0; i < 101; i++ { for i := 0; i < 101; i++ {
plugin.Command("cmd"+strconv.Itoa(i), exec) plugin.Command("cmd"+strconv.Itoa(i), exec)
} }
@@ -66,7 +66,7 @@ func TestAutoGenerateCommandsChecksLimitBeforeDelete(t *testing.T) {
func TestGatherCommandsForPluginReturnsSortedCommands(t *testing.T) { func TestGatherCommandsForPluginReturnsSortedCommands(t *testing.T) {
plugin := NewPlugin[NoData]("sorted") plugin := NewPlugin[NoData]("sorted")
exec := func(ctx *MsgContext, db NoData) error { return nil } exec := func(ctx *MessageContext, db NoData) error { return nil }
plugin.Command("zeta", exec) plugin.Command("zeta", exec)
plugin.Command("alpha", exec) plugin.Command("alpha", exec)
+13 -9
View File
@@ -85,7 +85,7 @@ func (c CommandArg) SetRequired() CommandArg {
// CommandExecutor is the function type that executes a command. // CommandExecutor is the function type that executes a command.
// It receives the message context and injected application data. // It receives the message context and injected application data.
// Returning a non-nil error routes it through the bot's error handler. // Returning a non-nil error routes it through the bot's error handler.
type CommandExecutor[T AppData] func(ctx *MsgContext, dbContext T) error type CommandExecutor[T AppData] func(ctx *MessageContext, dbContext T) error
// Command represents a bot command with arguments, description, and executor. // Command represents a bot command with arguments, description, and executor.
// Can be registered in a Plugin and optionally skipped from auto-generation. // Can be registered in a Plugin and optionally skipped from auto-generation.
@@ -98,18 +98,22 @@ type Command[T AppData] struct {
skipAutoCmd bool // If true, this command won't be auto-added to help menus skipAutoCmd bool // If true, this command won't be auto-added to help menus
} }
// NewCommand creates a new Command with the given command string, executor, and arguments. // NewCommand creates a new Command with the given identifier, executor, and arguments.
// The command string should not include the leading slash (e.g., "start", not "/start"). //
// The identifier is used as the routing key for both /-prefixed commands and
// callback payloads — the difference is registration: pass the result to
// Plugin.AddCommand/Plugin.Command for message routing, or to
// Plugin.AddPayload/Plugin.Payload for callback_data routing.
//
// For /-commands the identifier must not include the leading slash
// (e.g. "start", not "/start") and should match [_a-z0-9]{1,32} to satisfy
// Telegram's BotCommand validation. Payload identifiers may use any bytes
// that fit Telegram's callback_data limit, though the configured payload
// encoding may impose its own restrictions.
func NewCommand[T any](command string, exec CommandExecutor[T], args ...CommandArg) *Command[T] { func NewCommand[T any](command string, exec CommandExecutor[T], args ...CommandArg) *Command[T] {
return &Command[T]{command, "", exec, args, make(extypes.Slice[Middleware[T]], 0), false} return &Command[T]{command, "", exec, args, make(extypes.Slice[Middleware[T]], 0), false}
} }
// NewPayload creates a new callback payload handler command.
// The command string can contain any symbols, but it is recommended to use only "_", "-", ".", a-z, A-Z, and 0-9.
func NewPayload[T any](command string, exec CommandExecutor[T], args ...CommandArg) *Command[T] {
return &Command[T]{command, "", exec, args, make(extypes.Slice[Middleware[T]], 0), false}
}
// Use adds a middleware to the command's execution chain. // Use adds a middleware to the command's execution chain.
// Middlewares are executed in the order they are added. // Middlewares are executed in the order they are added.
func (c *Command[T]) Use(m Middleware[T]) *Command[T] { func (c *Command[T]) Use(m Middleware[T]) *Command[T] {
+1 -1
View File
@@ -5,7 +5,7 @@ Core concepts:
- Bot manages Telegram API access, update processing, logging, rate limiting, and dependency injection. - Bot manages Telegram API access, update processing, logging, rate limiting, and dependency injection.
- Plugins group commands, payloads, and non-command update handlers behind shared middleware. - Plugins group commands, payloads, and non-command update handlers behind shared middleware.
- MsgContext provides access to the current update and reply/edit/delete helpers. - MessageContext provides access to the current update and reply/edit/delete helpers.
- InlineKeyboard builds callback-driven keyboards and structured payloads. - InlineKeyboard builds callback-driven keyboards and structured payloads.
- DraftProvider accumulates multi-step replies before sending them. - DraftProvider accumulates multi-step replies before sending them.
- L10n stores key-based translations with fallback behavior. - L10n stores key-based translations with fallback behavior.
+15 -5
View File
@@ -14,8 +14,11 @@ type draftIDGenerator interface {
Next() uint64 Next() uint64
} }
// RandomDraftIDGenerator generates draft IDs using cryptographically secure random numbers. // RandomDraftIDGenerator generates draft IDs using math/rand/v2.
// Suitable for distributed systems or when ID predictability is undesirable. //
// Suitable for general use thanks to the wide 64-bit value space. Not suitable
// for security-sensitive purposes — use crypto/rand if unpredictability against
// an adversary matters.
type RandomDraftIDGenerator struct{} type RandomDraftIDGenerator struct{}
// Next returns a random 64-bit unsigned integer. // Next returns a random 64-bit unsigned integer.
@@ -29,7 +32,7 @@ type LinearDraftIDGenerator struct {
lastID atomic.Uint64 lastID atomic.Uint64
} }
// Next returns the next linear ID, atomically incremented.о // Next returns the next linear ID, atomically incremented.
func (g *LinearDraftIDGenerator) Next() uint64 { func (g *LinearDraftIDGenerator) Next() uint64 {
return g.lastID.Add(1) return g.lastID.Add(1)
} }
@@ -239,14 +242,21 @@ func (d *Draft) Flush() error {
} }
// Internal helper for Push that updates the server-side draft. // Internal helper for Push that updates the server-side draft.
//
// The candidate Message (current content + new text) is validated before any
// mutation, so a validation failure leaves the draft unchanged. After the
// validation passes, Message is committed locally regardless of whether the
// API call succeeds (per the Push docs: local state reflects the user's
// intent, network failures can be retried).
func (d *Draft) push(text string) error { func (d *Draft) push(text string) error {
if d.chatID == 0 { if d.chatID == 0 {
return ErrDraftChatIDZero return ErrDraftChatIDZero
} }
d.Message += text candidate := d.Message + text
if err := validateMessageText(d.Message); err != nil { if err := validateMessageText(candidate); err != nil {
return err return err
} }
d.Message = candidate
params := tgapi.SendMessageDraft{ params := tgapi.SendMessageDraft{
ChatID: d.chatID, ChatID: d.chatID,
DraftID: d.ID, DraftID: d.ID,
+19 -1
View File
@@ -19,7 +19,7 @@ func TestDraftFlushRequiresChatID(t *testing.T) {
} }
func TestMsgContextNewDraftWorksWithoutLimiter(t *testing.T) { func TestMsgContextNewDraftWorksWithoutLimiter(t *testing.T) {
ctx := &MsgContext{ ctx := &MessageContext{
API: &tgapi.API{}, API: &tgapi.API{},
Msg: &tgapi.Message{ Msg: &tgapi.Message{
Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate},
@@ -31,6 +31,7 @@ func TestMsgContextNewDraftWorksWithoutLimiter(t *testing.T) {
draft := ctx.NewDraft() draft := ctx.NewDraft()
if draft == nil { if draft == nil {
t.Fatal("expected draft") t.Fatal("expected draft")
return
} }
if draft.chatID != 42 { if draft.chatID != 42 {
t.Fatalf("unexpected chat id: %d", draft.chatID) t.Fatalf("unexpected chat id: %d", draft.chatID)
@@ -53,3 +54,20 @@ func TestDraftPushRejectsLongMessage(t *testing.T) {
t.Fatalf("expected ErrMessageTooLong, got %v", err) t.Fatalf("expected ErrMessageTooLong, got %v", err)
} }
} }
// TestDraftPushLeavesMessageUnchangedOnValidationFailure covers the validation
// order fix: when the candidate Message (current + new text) overflows the
// Telegram limit, the existing Message must remain intact so callers can
// recover and retry with a shorter payload instead of finding the draft in
// a half-mutated state.
func TestDraftPushLeavesMessageUnchangedOnValidationFailure(t *testing.T) {
draft := NewRandomDraftProvider(&tgapi.API{}).NewDraft(tgapi.ParseNone).SetChat(42, 0)
draft.Message = "hello"
if err := draft.Push(strings.Repeat("a", maxMessageTextLen+1)); !errors.Is(err, ErrMessageTooLong) {
t.Fatalf("expected ErrMessageTooLong, got %v", err)
}
if draft.Message != "hello" {
t.Fatalf("expected draft Message to stay %q, got %q", "hello", draft.Message)
}
}
+84 -11
View File
@@ -26,7 +26,7 @@ func (bot *Bot[T]) handle(parentCtx context.Context, u *tgapi.Update) {
ctx, cancel := context.WithCancel(parentCtx) ctx, cancel := context.WithCancel(parentCtx)
defer cancel() defer cancel()
msgCtx := &MsgContext{ msgCtx := &MessageContext{
Update: *u, API: bot.api, Update: *u, API: bot.api,
Logger: bot.logger, Logger: bot.logger,
errorTemplate: bot.errorTemplate, errorTemplate: bot.errorTemplate,
@@ -35,6 +35,7 @@ func (bot *Bot[T]) handle(parentCtx context.Context, u *tgapi.Update) {
sceneRuntime: bot, sceneRuntime: bot,
observer: bot.observer, observer: bot.observer,
payloadType: bot.payloadType, payloadType: bot.payloadType,
botID: bot.userID,
ctx: ctx, ctx: ctx,
} }
bot.prepareUpdateCtx(u, msgCtx) bot.prepareUpdateCtx(u, msgCtx)
@@ -114,7 +115,7 @@ func (bot *Bot[T]) handle(parentCtx context.Context, u *tgapi.Update) {
}) })
} }
func cloneMsgContext(src *MsgContext) *MsgContext { func cloneMsgContext(src *MessageContext) *MessageContext {
cloned := *src cloned := *src
if src.Args != nil { if src.Args != nil {
cloned.Args = append([]string(nil), src.Args...) cloned.Args = append([]string(nil), src.Args...)
@@ -154,20 +155,92 @@ func decodeBase64Payload(s string) (CallbackData, error) {
return decodeJSONPayload(string(b)) return decodeJSONPayload(string(b))
} }
func encodeCompactPayload(d CallbackData) (string, error) { // Compact payload format: cmd|arg1,arg2,...
args := strings.Join(d.Args, ",") // Bytes \, |, and , inside a part are escaped with a leading backslash so the
return d.Command + "|" + args, nil // payload round-trips without ambiguity. Encoding/decoding operate byte-wise
// because all separators are single-byte ASCII; multi-byte UTF-8 code points
// pass through unchanged.
func encodeCompactPart(s string) string {
if !strings.ContainsAny(s, `\|,`) {
return s
}
var b strings.Builder
b.Grow(len(s) + 2)
for i := 0; i < len(s); i++ {
switch s[i] {
case '\\', '|', ',':
b.WriteByte('\\')
}
b.WriteByte(s[i])
}
return b.String()
} }
func decodeCompactPart(s string) string {
if !strings.Contains(s, `\`) {
return s
}
var b strings.Builder
b.Grow(len(s))
for i := 0; i < len(s); i++ {
if s[i] == '\\' && i+1 < len(s) {
b.WriteByte(s[i+1])
i++
continue
}
b.WriteByte(s[i])
}
return b.String()
}
func encodeCompactPayload(d CallbackData) (string, error) {
var b strings.Builder
b.WriteString(encodeCompactPart(d.Command))
b.WriteByte('|')
for i, a := range d.Args {
if i > 0 {
b.WriteByte(',')
}
b.WriteString(encodeCompactPart(a))
}
return b.String(), nil
}
func decodeCompactPayload(s string) (CallbackData, error) { func decodeCompactPayload(s string) (CallbackData, error) {
values := strings.SplitN(s, "|", 2) sepIdx := -1
if len(values) != 2 { for i := 0; i < len(s); i++ {
if s[i] == '\\' && i+1 < len(s) {
i++
continue
}
if s[i] == '|' {
sepIdx = i
break
}
}
if sepIdx == -1 {
return CallbackData{}, errors.New("invalid payload") return CallbackData{}, errors.New("invalid payload")
} }
cmd, argsRaw := values[0], values[1] cmd := decodeCompactPart(s[:sepIdx])
var args []string argsRaw := s[sepIdx+1:]
if argsRaw != "" { if argsRaw == "" {
args = strings.Split(argsRaw, ",") return CallbackData{Command: cmd}, nil
} }
var args []string
start := 0
for i := 0; i < len(argsRaw); i++ {
if argsRaw[i] == '\\' && i+1 < len(argsRaw) {
i++
continue
}
if argsRaw[i] == ',' {
args = append(args, decodeCompactPart(argsRaw[start:i]))
start = i + 1
}
}
args = append(args, decodeCompactPart(argsRaw[start:]))
return CallbackData{Command: cmd, Args: args}, nil return CallbackData{Command: cmd, Args: args}, nil
} }
func encodeCompactBase64Payload(d CallbackData) (string, error) { func encodeCompactBase64Payload(d CallbackData) (string, error) {
+26 -26
View File
@@ -64,7 +64,7 @@ func TestBotMiddlewareReceivesLogger(t *testing.T) {
bot := &Bot[NoData]{ bot := &Bot[NoData]{
logger: logger, logger: logger,
middlewares: []Middleware[NoData]{ middlewares: []Middleware[NoData]{
NewMiddleware("logger-check", func(ctx *MsgContext, db NoData) bool { NewMiddleware("logger-check", func(ctx *MessageContext, db NoData) bool {
called = true called = true
if ctx.Logger != logger { if ctx.Logger != logger {
t.Fatalf("expected bot logger in middleware context, got %#v", ctx.Logger) t.Fatalf("expected bot logger in middleware context, got %#v", ctx.Logger)
@@ -90,7 +90,7 @@ func TestBotMiddlewareReceivesLogger(t *testing.T) {
func TestAddUpdateHandlerRejectsReservedUpdateTypes(t *testing.T) { func TestAddUpdateHandlerRejectsReservedUpdateTypes(t *testing.T) {
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
handler := func(ctx *MsgContext, db NoData) error { return nil } handler := func(ctx *MessageContext, db NoData) error { return nil }
for _, updateType := range []tgapi.UpdateType{ for _, updateType := range []tgapi.UpdateType{
tgapi.UpdateTypeMessage, tgapi.UpdateTypeMessage,
@@ -376,7 +376,7 @@ func TestPrepareUpdateCtxContract(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
bot := &Bot[NoData]{} bot := &Bot[NoData]{}
ctx := &MsgContext{} ctx := &MessageContext{}
bot.prepareUpdateCtx(tt.update, ctx) bot.prepareUpdateCtx(tt.update, ctx)
if got := ctx.Msg != nil; got != tt.wantMsg { if got := ctx.Msg != nil; got != tt.wantMsg {
@@ -450,7 +450,7 @@ func TestHandleUpdateHandlersPopulateFromContext(t *testing.T) {
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
called := false called := false
plugin := NewPlugin[NoData]("test").AddUpdateHandler(tt.update.Type, func(ctx *MsgContext, db NoData) error { plugin := NewPlugin[NoData]("test").AddUpdateHandler(tt.update.Type, func(ctx *MessageContext, db NoData) error {
called = true called = true
if ctx.Update.UpdateID != tt.update.UpdateID { if ctx.Update.UpdateID != tt.update.UpdateID {
t.Fatalf("unexpected update in context: got %d want %d", ctx.Update.UpdateID, tt.update.UpdateID) t.Fatalf("unexpected update in context: got %d want %d", ctx.Update.UpdateID, tt.update.UpdateID)
@@ -488,7 +488,7 @@ func TestHandleUpdateHandlersReceiveIsolatedContexts(t *testing.T) {
firstCalled := false firstCalled := false
secondCalled := false secondCalled := false
first := NewPlugin[NoData]("first").AddUpdateHandler(tgapi.UpdateTypeInlineQuery, func(ctx *MsgContext, db NoData) error { first := NewPlugin[NoData]("first").AddUpdateHandler(tgapi.UpdateTypeInlineQuery, func(ctx *MessageContext, db NoData) error {
firstCalled = true firstCalled = true
if ctx.FromID != 41 { if ctx.FromID != 41 {
t.Fatalf("unexpected FromID in first handler: got %d want 41", ctx.FromID) t.Fatalf("unexpected FromID in first handler: got %d want 41", ctx.FromID)
@@ -499,7 +499,7 @@ func TestHandleUpdateHandlersReceiveIsolatedContexts(t *testing.T) {
ctx.Args = []string{"mutated"} ctx.Args = []string{"mutated"}
return nil return nil
}) })
second := NewPlugin[NoData]("second").AddUpdateHandler(tgapi.UpdateTypeInlineQuery, func(ctx *MsgContext, db NoData) error { second := NewPlugin[NoData]("second").AddUpdateHandler(tgapi.UpdateTypeInlineQuery, func(ctx *MessageContext, db NoData) error {
secondCalled = true secondCalled = true
if ctx.From == nil { if ctx.From == nil {
t.Fatal("expected ctx.From to remain populated for second handler") t.Fatal("expected ctx.From to remain populated for second handler")
@@ -541,7 +541,7 @@ func TestHandleUpdateHandlersReceiveIsolatedContexts(t *testing.T) {
func TestHandleUpdateObserverEmitsUpdateErrors(t *testing.T) { func TestHandleUpdateObserverEmitsUpdateErrors(t *testing.T) {
observer := &recordingObserver{} observer := &recordingObserver{}
plugin := NewPlugin[NoData]("test").AddUpdateHandler(tgapi.UpdateTypeInlineQuery, func(ctx *MsgContext, db NoData) error { plugin := NewPlugin[NoData]("test").AddUpdateHandler(tgapi.UpdateTypeInlineQuery, func(ctx *MessageContext, db NoData) error {
return AsUserError(errors.New("update failed")) return AsUserError(errors.New("update failed"))
}) })
@@ -596,7 +596,7 @@ func TestHandleObserverCompletesUpdateWhenBotMiddlewareBlocks(t *testing.T) {
logger: sneklog.NewLogger(), logger: sneklog.NewLogger(),
observer: observer, observer: observer,
middlewares: []Middleware[NoData]{ middlewares: []Middleware[NoData]{
NewMiddleware("block", func(ctx *MsgContext, db NoData) bool { NewMiddleware("block", func(ctx *MessageContext, db NoData) bool {
return false return false
}), }),
}, },
@@ -632,7 +632,7 @@ func TestHandleMessageFallbackRunsAfterCommandMiss(t *testing.T) {
observer := &recordingObserver{} observer := &recordingObserver{}
called := false called := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.SetMessageFallback(func(ctx *MsgContext, db NoData) error { plugin.SetMessageFallback(func(ctx *MessageContext, db NoData) error {
called = true called = true
if ctx.Text != "/missing hello world" { if ctx.Text != "/missing hello world" {
t.Fatalf("unexpected fallback text: got %q", ctx.Text) t.Fatalf("unexpected fallback text: got %q", ctx.Text)
@@ -687,7 +687,7 @@ func TestHandleMessageFallbackRunsAfterCommandMiss(t *testing.T) {
func TestHandleMessageFallbackRunsForPlainText(t *testing.T) { func TestHandleMessageFallbackRunsForPlainText(t *testing.T) {
called := false called := false
plugin := NewPlugin[NoData]("test").SetMessageFallback(func(ctx *MsgContext, db NoData) error { plugin := NewPlugin[NoData]("test").SetMessageFallback(func(ctx *MessageContext, db NoData) error {
called = true called = true
if ctx.Text != "hello fallback" { if ctx.Text != "hello fallback" {
t.Fatalf("unexpected fallback text: got %q", ctx.Text) t.Fatalf("unexpected fallback text: got %q", ctx.Text)
@@ -723,10 +723,10 @@ func TestHandleMessageFallbackRunsForPlainText(t *testing.T) {
func TestHandleMessageFallbackRespectsMiddleware(t *testing.T) { func TestHandleMessageFallbackRespectsMiddleware(t *testing.T) {
called := false called := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.AddMiddleware(NewMiddleware("block", func(ctx *MsgContext, db NoData) bool { plugin.AddMiddleware(NewMiddleware("block", func(ctx *MessageContext, db NoData) bool {
return false return false
})) }))
plugin.SetMessageFallback(func(ctx *MsgContext, db NoData) error { plugin.SetMessageFallback(func(ctx *MessageContext, db NoData) error {
called = true called = true
return nil return nil
}) })
@@ -757,11 +757,11 @@ func TestHandleMessageFallbackDoesNotRunWhenCommandMatches(t *testing.T) {
commandCalled := false commandCalled := false
fallbackCalled := false fallbackCalled := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Command("start", func(ctx *MsgContext, db NoData) error { plugin.Command("start", func(ctx *MessageContext, db NoData) error {
commandCalled = true commandCalled = true
return nil return nil
}) })
plugin.SetMessageFallback(func(ctx *MsgContext, db NoData) error { plugin.SetMessageFallback(func(ctx *MessageContext, db NoData) error {
fallbackCalled = true fallbackCalled = true
return nil return nil
}) })
@@ -794,7 +794,7 @@ func TestHandleMessageFallbackDoesNotRunWhenCommandMatches(t *testing.T) {
func TestHandleChannelPostCommandWithSenderChat(t *testing.T) { func TestHandleChannelPostCommandWithSenderChat(t *testing.T) {
called := false called := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Command("ping", func(ctx *MsgContext, db NoData) error { plugin.Command("ping", func(ctx *MessageContext, db NoData) error {
called = true called = true
if ctx.Msg == nil { if ctx.Msg == nil {
t.Fatal("expected message context") t.Fatal("expected message context")
@@ -841,7 +841,7 @@ func TestCommandHandlerBindArgsEndToEnd(t *testing.T) {
var got banInput var got banInput
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Command("ban", func(ctx *MsgContext, db NoData) error { plugin.Command("ban", func(ctx *MessageContext, db NoData) error {
return ctx.BindArgs(&got) return ctx.BindArgs(&got)
}, },
NewCommandArg("user_id").SetValueType(CommandValueInt).SetRequired(), NewCommandArg("user_id").SetValueType(CommandValueInt).SetRequired(),
@@ -878,7 +878,7 @@ func TestPayloadHandlerBindArgsEndToEnd(t *testing.T) {
var got payloadInput var got payloadInput
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Payload("approve", func(ctx *MsgContext, db NoData) error { plugin.Payload("approve", func(ctx *MessageContext, db NoData) error {
return ctx.BindArgs(&got) return ctx.BindArgs(&got)
}, },
NewCommandArg("id").SetValueType(CommandValueInt).SetRequired(), NewCommandArg("id").SetValueType(CommandValueInt).SetRequired(),
@@ -920,11 +920,11 @@ func TestHandleEditedMessageStaysOutOfCommandFlow(t *testing.T) {
updateCalled := false updateCalled := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Command("ping", func(ctx *MsgContext, db NoData) error { plugin.Command("ping", func(ctx *MessageContext, db NoData) error {
commandCalled = true commandCalled = true
return nil return nil
}) })
plugin.AddUpdateHandler(tgapi.UpdateTypeEditedMessage, func(ctx *MsgContext, db NoData) error { plugin.AddUpdateHandler(tgapi.UpdateTypeEditedMessage, func(ctx *MessageContext, db NoData) error {
updateCalled = true updateCalled = true
if ctx.Msg == nil { if ctx.Msg == nil {
t.Fatal("expected ctx.Msg in edited message handler") t.Fatal("expected ctx.Msg in edited message handler")
@@ -968,11 +968,11 @@ func TestHandleEditedChannelPostStaysOutOfCommandFlow(t *testing.T) {
updateCalled := false updateCalled := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Command("ping", func(ctx *MsgContext, db NoData) error { plugin.Command("ping", func(ctx *MessageContext, db NoData) error {
commandCalled = true commandCalled = true
return nil return nil
}) })
plugin.AddUpdateHandler(tgapi.UpdateTypeEditedChannelPost, func(ctx *MsgContext, db NoData) error { plugin.AddUpdateHandler(tgapi.UpdateTypeEditedChannelPost, func(ctx *MessageContext, db NoData) error {
updateCalled = true updateCalled = true
if ctx.Msg == nil { if ctx.Msg == nil {
t.Fatal("expected ctx.Msg in edited channel post handler") t.Fatal("expected ctx.Msg in edited channel post handler")
@@ -1007,7 +1007,7 @@ func TestHandleEditedChannelPostStaysOutOfCommandFlow(t *testing.T) {
func TestHandleCallbackPopulatesMessageTargets(t *testing.T) { func TestHandleCallbackPopulatesMessageTargets(t *testing.T) {
called := false called := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Payload("approve", func(ctx *MsgContext, db NoData) error { plugin.Payload("approve", func(ctx *MessageContext, db NoData) error {
called = true called = true
if ctx.CallbackQueryID != "cb-msg" { if ctx.CallbackQueryID != "cb-msg" {
t.Fatalf("unexpected CallbackQueryID: %q", ctx.CallbackQueryID) t.Fatalf("unexpected CallbackQueryID: %q", ctx.CallbackQueryID)
@@ -1066,7 +1066,7 @@ func TestHandleCallbackPopulatesMessageTargets(t *testing.T) {
func TestHandleCallbackPopulatesInlineTargets(t *testing.T) { func TestHandleCallbackPopulatesInlineTargets(t *testing.T) {
called := false called := false
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Payload("inline.approve", func(ctx *MsgContext, db NoData) error { plugin.Payload("inline.approve", func(ctx *MessageContext, db NoData) error {
called = true called = true
if ctx.CallbackQueryID != "cb-inline" { if ctx.CallbackQueryID != "cb-inline" {
t.Fatalf("unexpected CallbackQueryID: %q", ctx.CallbackQueryID) t.Fatalf("unexpected CallbackQueryID: %q", ctx.CallbackQueryID)
@@ -1122,7 +1122,7 @@ func TestHandleCallbackPopulatesInlineTargets(t *testing.T) {
func TestHandleCallbackObserverEmitsPayloadEvents(t *testing.T) { func TestHandleCallbackObserverEmitsPayloadEvents(t *testing.T) {
observer := &recordingObserver{} observer := &recordingObserver{}
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
plugin.Payload("approve", func(ctx *MsgContext, db NoData) error { plugin.Payload("approve", func(ctx *MessageContext, db NoData) error {
return nil return nil
}) })
@@ -1173,7 +1173,7 @@ func TestHandleCallbackObserverEmitsPayloadErrors(t *testing.T) {
observer := &recordingObserver{} observer := &recordingObserver{}
plugin := NewPlugin[NoData]("test") plugin := NewPlugin[NoData]("test")
wantErr := AsInternalError(errors.New("boom")) wantErr := AsInternalError(errors.New("boom"))
plugin.Payload("approve", func(ctx *MsgContext, db NoData) error { plugin.Payload("approve", func(ctx *MessageContext, db NoData) error {
return wantErr return wantErr
}) })
@@ -1236,7 +1236,7 @@ func TestHandleCallbackObserverEmitsDecodeErrors(t *testing.T) {
Data: "{not-json", Data: "{not-json",
From: tgapi.User{ID: 7}, From: tgapi.User{ID: 7},
}, },
}, &MsgContext{ }, &MessageContext{
Update: tgapi.Update{ Update: tgapi.Update{
UpdateID: 34, UpdateID: 34,
Type: tgapi.UpdateTypeCallbackQuery, Type: tgapi.UpdateTypeCallbackQuery,
+19 -11
View File
@@ -203,6 +203,8 @@ func (in *InlineKeyboard) SetMaxRow(maxRow int) *InlineKeyboard {
return in return in
} }
func (in *InlineKeyboard) GetMaxRow() int { return in.maxRow }
// Internal helper that appends a button and auto-flushes a full row. // Internal helper that appends a button and auto-flushes a full row.
func (in *InlineKeyboard) append(button tgapi.InlineKeyboardButton) *InlineKeyboard { func (in *InlineKeyboard) append(button tgapi.InlineKeyboardButton) *InlineKeyboard {
if in.CurrentLine.Len() == in.maxRow { if in.CurrentLine.Len() == in.maxRow {
@@ -302,18 +304,18 @@ func NewCallbackData(command string, args ...any) CallbackData {
} }
} }
// All To* encoders return an empty string when serialization fails. Telegram
// rejects empty callback_data, so an empty result surfaces a real bug rather
// than masking it with a stub payload that silently routes to no handler.
// Build CallbackData from primitives (string, []string) only — the encoders
// have no failure modes for that input.
// ToJSON serializes the CallbackData to a JSON string. // ToJSON serializes the CallbackData to a JSON string.
// // Returns an empty string if serialization fails.
// If serialization fails (e.g., due to unmarshalable fields), returns a fallback
// JSON object: {"cmd":""} to prevent breaking Telegram's API.
//
// This fallback ensures the bot receives a valid JSON payload even if internal
// errors occur — avoiding "invalid callback_data" errors from Telegram.
func (d CallbackData) ToJSON() string { func (d CallbackData) ToJSON() string {
data, err := encodeJSONPayload(d) data, err := encodeJSONPayload(d)
if err != nil { if err != nil {
// Fallback: return minimal valid JSON to avoid Telegram API rejection return ""
return `{"cmd":""}`
} }
return data return data
} }
@@ -323,25 +325,31 @@ func (d CallbackData) ToJSON() string {
func (d CallbackData) ToBase64() string { func (d CallbackData) ToBase64() string {
data, err := encodeBase64Payload(d) data, err := encodeBase64Payload(d)
if err != nil { if err != nil {
return `` return ""
} }
return data return data
} }
// ToCompact serializes the CallbackData to a compact delimited string. // ToCompact serializes the CallbackData to a compact delimited string.
// Returns an empty string if serialization fails.
//
// The compact format coalesces "no args" with "single empty arg" — both
// produce "cmd|" and decode back to nil args. Use ToJSON or ToBase64 when
// that distinction must be preserved.
func (d CallbackData) ToCompact() string { func (d CallbackData) ToCompact() string {
data, err := encodeCompactPayload(d) data, err := encodeCompactPayload(d)
if err != nil { if err != nil {
return `` return ""
} }
return data return data
} }
// ToCompactBase64 serializes the CallbackData to compact text and then encodes it as Base64. // ToCompactBase64 serializes the CallbackData to compact text and then encodes it as Base64.
// Returns an empty string if serialization or encoding fails.
func (d CallbackData) ToCompactBase64() string { func (d CallbackData) ToCompactBase64() string {
data, err := encodeCompactBase64Payload(d) data, err := encodeCompactBase64Payload(d)
if err != nil { if err != nil {
return `` return ""
} }
return data return data
} }
+54
View File
@@ -150,6 +150,60 @@ func TestDecodePayloadAcceptsCompactBase64KeyboardPayloadWhenBotPrefersJSON(t *t
} }
} }
// TestCompactPayloadRoundTripsWithSeparatorChars guards the compact-encoding
// escape fix. Args containing the , | or \ separator bytes previously corrupted
// on decode; now they must round-trip exactly.
//
// Note: the compact format coalesces "no args" with "single empty arg" — both
// emit "cmd|" and decode to nil args. Use other encodings if that distinction
// matters.
func TestCompactPayloadRoundTripsWithSeparatorChars(t *testing.T) {
tests := []struct {
name string
data CallbackData
}{
{name: "plain", data: CallbackData{Command: "cmd", Args: []string{"one", "two"}}},
{name: "no args", data: CallbackData{Command: "cmd"}},
{name: "comma in arg", data: CallbackData{Command: "cmd", Args: []string{"a,b", "c"}}},
{name: "pipe in arg", data: CallbackData{Command: "cmd", Args: []string{"a|b", "c"}}},
{name: "backslash in arg", data: CallbackData{Command: "cmd", Args: []string{`a\b`, "c"}}},
{name: "all specials in arg", data: CallbackData{Command: "cmd", Args: []string{`a,b|c\d`}}},
{name: "specials in command", data: CallbackData{Command: "a|b,c", Args: []string{"x"}}},
{name: "two empty args", data: CallbackData{Command: "cmd", Args: []string{"", ""}}},
{name: "utf8 args", data: CallbackData{Command: "cmd", Args: []string{"привет", "мир"}}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
encoded, err := encodeCompactPayload(tt.data)
if err != nil {
t.Fatalf("encodeCompactPayload returned error: %v", err)
}
got, err := decodeCompactPayload(encoded)
if err != nil {
t.Fatalf("decodeCompactPayload returned error: %v", err)
}
if got.Command != tt.data.Command {
t.Fatalf("command mismatch: got %q want %q (encoded=%q)", got.Command, tt.data.Command, encoded)
}
if len(got.Args) != len(tt.data.Args) {
t.Fatalf("args length mismatch: got %v want %v (encoded=%q)", got.Args, tt.data.Args, encoded)
}
for i := range tt.data.Args {
if got.Args[i] != tt.data.Args[i] {
t.Fatalf("arg %d mismatch: got %q want %q (encoded=%q)", i, got.Args[i], tt.data.Args[i], encoded)
}
}
})
}
}
func TestCompactPayloadDecodeRejectsMissingSeparator(t *testing.T) {
if _, err := decodeCompactPayload("noseparator"); err == nil {
t.Fatal("expected error decoding payload without separator")
}
}
func TestDecodePayloadStrictRejectsCompactMismatchedType(t *testing.T) { func TestDecodePayloadStrictRejectsCompactMismatchedType(t *testing.T) {
kb := NewInlineKeyboardCompact(1). kb := NewInlineKeyboardCompact(1).
AddCallbackButton("A", "cmd", 1) AddCallbackButton("A", "cmd", 1)
+2 -1
View File
@@ -43,9 +43,10 @@ import (
// } // }
func (bot *Bot[T]) Updates(ctx context.Context) ([]tgapi.Update, error) { func (bot *Bot[T]) Updates(ctx context.Context) ([]tgapi.Update, error) {
offset := bot.GetUpdateOffset() offset := bot.GetUpdateOffset()
timeout := bot.pollTimeout
params := tgapi.UpdateParams{ params := tgapi.UpdateParams{
Offset: new(offset), Offset: new(offset),
Timeout: new(30), Timeout: new(timeout),
AllowedUpdates: bot.GetUpdateTypes(), AllowedUpdates: bot.GetUpdateTypes(),
} }
+63 -62
View File
@@ -14,10 +14,10 @@ import (
"git.scuroneko.dev/scuroneko/sneklog/v2" "git.scuroneko.dev/scuroneko/sneklog/v2"
) )
// MsgContext holds the normalized per-update context passed to command, payload, // MessageContext holds the normalized per-update context passed to command, payload,
// scene, middleware, and generic update handlers. // scene, middleware, and generic update handlers.
// //
// MsgContext is populated from the current Telegram update before handler routing. // MessageContext is populated from the current Telegram update before handler routing.
// Not every field is guaranteed for every update kind. In particular: // Not every field is guaranteed for every update kind. In particular:
// - Update is always present. // - Update is always present.
// - Msg is populated only for update kinds that carry a Telegram message object. // - Msg is populated only for update kinds that carry a Telegram message object.
@@ -27,10 +27,10 @@ import (
// - CallbackQueryID, CallbackMsgID, and InlineMsgID are populated only for // - CallbackQueryID, CallbackMsgID, and InlineMsgID are populated only for
// callback query handling when the corresponding callback targets exist. // callback query handling when the corresponding callback targets exist.
// //
// Helper methods on MsgContext may require a message-backed context. For example, // Helper methods on MessageContext may require a message-backed context. For example,
// reply helpers need Msg, while inline callback edit helpers can work through // reply helpers need Msg, while inline callback edit helpers can work through
// InlineMsgID when there is no chat message. // InlineMsgID when there is no chat message.
type MsgContext struct { type MessageContext struct {
API *tgapi.API API *tgapi.API
Update tgapi.Update Update tgapi.Update
@@ -81,21 +81,22 @@ type MsgContext struct {
payloadType BotPayloadType payloadType BotPayloadType
sceneRuntime sceneRuntime sceneRuntime sceneRuntime
observer Observer observer Observer
botID int64
ctx context.Context ctx context.Context
} }
// AnswerMessage represents a message sent or edited via MsgContext. // AnswerMessage represents a message sent or edited via MessageContext.
// It holds metadata to allow further editing or deletion. // It holds metadata to allow further editing or deletion.
type AnswerMessage struct { type AnswerMessage struct {
MessageID int MessageID int
Text string Text string
IsMedia bool IsMedia bool
ctx *MsgContext // internal back-reference ctx *MessageContext // internal back-reference
} }
// Internal helper for text edits with optional keyboard and parse mode. // Internal helper for text edits with optional keyboard and parse mode.
func (ctx *MsgContext) edit(messageID int, text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { func (ctx *MessageContext) edit(messageID int, text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage {
if err := validateMessageText(text); err != nil { if err := validateMessageText(text); err != nil {
ctx.Logger.Errorln(err) ctx.Logger.Errorln(err)
return nil return nil
@@ -146,7 +147,7 @@ func (m *AnswerMessage) EditMarkdown(text string) *AnswerMessage {
} }
// Internal helper for editing callback-linked messages. // Internal helper for editing callback-linked messages.
func (ctx *MsgContext) editCallback(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { func (ctx *MessageContext) editCallback(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage {
if ctx.CallbackMsgID == 0 && ctx.InlineMsgID == "" { if ctx.CallbackMsgID == 0 && ctx.InlineMsgID == "" {
ctx.Logger.Errorln(ErrCallbackMessageMissing) ctx.Logger.Errorln(ErrCallbackMessageMissing)
return nil return nil
@@ -155,31 +156,31 @@ func (ctx *MsgContext) editCallback(text string, keyboard *InlineKeyboard, parse
} }
// EditCallback edits the callback message using plain text (ParseNone). // EditCallback edits the callback message using plain text (ParseNone).
func (ctx *MsgContext) EditCallback(text string, keyboard *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) EditCallback(text string, keyboard *InlineKeyboard) *AnswerMessage {
return ctx.editCallback(text, keyboard, tgapi.ParseNone) return ctx.editCallback(text, keyboard, tgapi.ParseNone)
} }
// EditCallbackMarkdown edits the callback message using MarkdownV2. // EditCallbackMarkdown edits the callback message using MarkdownV2.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) EditCallbackMarkdown(text string, keyboard *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) EditCallbackMarkdown(text string, keyboard *InlineKeyboard) *AnswerMessage {
return ctx.editCallback(text, keyboard, tgapi.ParseMarkdownV2) return ctx.editCallback(text, keyboard, tgapi.ParseMarkdownV2)
} }
// EditCallbackf formats a string using fmt.Sprintf and edits the callback message with plain text. // EditCallbackf formats a string using fmt.Sprintf and edits the callback message with plain text.
func (ctx *MsgContext) EditCallbackf(format string, keyboard *InlineKeyboard, args ...any) *AnswerMessage { func (ctx *MessageContext) EditCallbackf(format string, keyboard *InlineKeyboard, args ...any) *AnswerMessage {
return ctx.editCallback(fmt.Sprintf(format, args...), keyboard, tgapi.ParseNone) return ctx.editCallback(fmt.Sprintf(format, args...), keyboard, tgapi.ParseNone)
} }
// EditCallbackfMarkdown formats a string using fmt.Sprintf and edits the callback message with MarkdownV2. // EditCallbackfMarkdown formats a string using fmt.Sprintf and edits the callback message with MarkdownV2.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) EditCallbackfMarkdown(format string, keyboard *InlineKeyboard, args ...any) *AnswerMessage { func (ctx *MessageContext) EditCallbackfMarkdown(format string, keyboard *InlineKeyboard, args ...any) *AnswerMessage {
return ctx.editCallback(fmt.Sprintf(format, args...), keyboard, tgapi.ParseMarkdownV2) return ctx.editCallback(fmt.Sprintf(format, args...), keyboard, tgapi.ParseMarkdownV2)
} }
// Internal helper for media-caption edits. // Internal helper for media-caption edits.
func (ctx *MsgContext) editPhotoText(messageID int, text string, kb *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { func (ctx *MessageContext) editPhotoText(messageID int, text string, kb *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage {
if err := validateCaptionText(text); err != nil { if err := validateCaptionText(text); err != nil {
ctx.Logger.Errorln(err) ctx.Logger.Errorln(err)
return nil return nil
@@ -241,7 +242,7 @@ func (m *AnswerMessage) EditCaptionKeyboardMarkdown(text string, kb *InlineKeybo
} }
// Internal helper for message replies with optional keyboard and parse mode. // Internal helper for message replies with optional keyboard and parse mode.
func (ctx *MsgContext) answer(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { func (ctx *MessageContext) answer(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage {
if ctx.Msg == nil { if ctx.Msg == nil {
ctx.Logger.Errorln(ErrMessageContextNil) ctx.Logger.Errorln(ErrMessageContextNil)
return nil return nil
@@ -276,7 +277,7 @@ func (ctx *MsgContext) answer(text string, keyboard *InlineKeyboard, parseMode t
} }
// Answer sends a plain text message (ParseNone). // Answer sends a plain text message (ParseNone).
func (ctx *MsgContext) Answer(text string) *AnswerMessage { func (ctx *MessageContext) Answer(text string) *AnswerMessage {
return ctx.answer(text, nil, tgapi.ParseNone) return ctx.answer(text, nil, tgapi.ParseNone)
} }
@@ -284,54 +285,54 @@ func (ctx *MsgContext) Answer(text string) *AnswerMessage {
// //
// The text is split into Telegram-safe chunks. Returned messages preserve send // The text is split into Telegram-safe chunks. Returned messages preserve send
// order. If a chunk fails to send, already-sent messages are returned. // order. If a chunk fails to send, already-sent messages are returned.
func (ctx *MsgContext) AnswerLong(text string) []*AnswerMessage { func (ctx *MessageContext) AnswerLong(text string) []*AnswerMessage {
return ctx.answerLong(text, nil, tgapi.ParseNone) return ctx.answerLong(text, nil, tgapi.ParseNone)
} }
// AnswerMarkdown sends a message using MarkdownV2 formatting. // AnswerMarkdown sends a message using MarkdownV2 formatting.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) AnswerMarkdown(text string) *AnswerMessage { func (ctx *MessageContext) AnswerMarkdown(text string) *AnswerMessage {
return ctx.answer(text, nil, tgapi.ParseMarkdownV2) return ctx.answer(text, nil, tgapi.ParseMarkdownV2)
} }
// Answerf formats a string using fmt.Sprintf and sends it as a plain text message. // Answerf formats a string using fmt.Sprintf and sends it as a plain text message.
func (ctx *MsgContext) Answerf(template string, args ...any) *AnswerMessage { func (ctx *MessageContext) Answerf(template string, args ...any) *AnswerMessage {
return ctx.answer(fmt.Sprintf(template, args...), nil, tgapi.ParseNone) return ctx.answer(fmt.Sprintf(template, args...), nil, tgapi.ParseNone)
} }
// AnswerLongf formats a string using fmt.Sprintf and sends it as one or more plain-text messages. // AnswerLongf formats a string using fmt.Sprintf and sends it as one or more plain-text messages.
func (ctx *MsgContext) AnswerLongf(template string, args ...any) []*AnswerMessage { func (ctx *MessageContext) AnswerLongf(template string, args ...any) []*AnswerMessage {
return ctx.answerLong(fmt.Sprintf(template, args...), nil, tgapi.ParseNone) return ctx.answerLong(fmt.Sprintf(template, args...), nil, tgapi.ParseNone)
} }
// AnswerfMarkdown formats a string using fmt.Sprintf and sends it using MarkdownV2. // AnswerfMarkdown formats a string using fmt.Sprintf and sends it using MarkdownV2.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) AnswerfMarkdown(template string, args ...any) *AnswerMessage { func (ctx *MessageContext) AnswerfMarkdown(template string, args ...any) *AnswerMessage {
return ctx.answer(fmt.Sprintf(template, args...), nil, tgapi.ParseMarkdownV2) return ctx.answer(fmt.Sprintf(template, args...), nil, tgapi.ParseMarkdownV2)
} }
// Keyboard sends a message with an inline keyboard (plain text). // Keyboard sends a message with an inline keyboard (plain text).
func (ctx *MsgContext) Keyboard(text string, kb *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) Keyboard(text string, kb *InlineKeyboard) *AnswerMessage {
return ctx.answer(text, kb, tgapi.ParseNone) return ctx.answer(text, kb, tgapi.ParseNone)
} }
// KeyboardLong sends long plain text split across multiple messages. // KeyboardLong sends long plain text split across multiple messages.
// //
// The inline keyboard is attached only to the final chunk. // The inline keyboard is attached only to the final chunk.
func (ctx *MsgContext) KeyboardLong(text string, kb *InlineKeyboard) []*AnswerMessage { func (ctx *MessageContext) KeyboardLong(text string, kb *InlineKeyboard) []*AnswerMessage {
return ctx.answerLong(text, kb, tgapi.ParseNone) return ctx.answerLong(text, kb, tgapi.ParseNone)
} }
// KeyboardMarkdown sends a message with an inline keyboard using MarkdownV2. // KeyboardMarkdown sends a message with an inline keyboard using MarkdownV2.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) KeyboardMarkdown(text string, keyboard *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) KeyboardMarkdown(text string, keyboard *InlineKeyboard) *AnswerMessage {
return ctx.answer(text, keyboard, tgapi.ParseMarkdownV2) return ctx.answer(text, keyboard, tgapi.ParseMarkdownV2)
} }
func (ctx *MsgContext) answerLong(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) []*AnswerMessage { func (ctx *MessageContext) answerLong(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) []*AnswerMessage {
if parseMode != tgapi.ParseNone { if parseMode != tgapi.ParseNone {
ctx.Logger.Errorln(ErrMessageSplitImpossible) ctx.Logger.Errorln(ErrMessageSplitImpossible)
return nil return nil
@@ -371,7 +372,7 @@ func (ctx *MsgContext) answerLong(text string, keyboard *InlineKeyboard, parseMo
} }
// Internal helper for photo replies with optional caption and keyboard. // Internal helper for photo replies with optional caption and keyboard.
func (ctx *MsgContext) answerPhoto(photoID, text string, kb *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { func (ctx *MessageContext) answerPhoto(photoID, text string, kb *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage {
if ctx.Msg == nil { if ctx.Msg == nil {
ctx.Logger.Errorln(ErrMessageContextNil) ctx.Logger.Errorln(ErrMessageContextNil)
return nil return nil
@@ -407,43 +408,43 @@ func (ctx *MsgContext) answerPhoto(photoID, text string, kb *InlineKeyboard, par
} }
// AnswerPhoto sends a photo with plain text caption. // AnswerPhoto sends a photo with plain text caption.
func (ctx *MsgContext) AnswerPhoto(photoID, text string) *AnswerMessage { func (ctx *MessageContext) AnswerPhoto(photoID, text string) *AnswerMessage {
return ctx.answerPhoto(photoID, text, nil, tgapi.ParseNone) return ctx.answerPhoto(photoID, text, nil, tgapi.ParseNone)
} }
// AnswerPhotoMarkdown sends a photo with MarkdownV2 caption. // AnswerPhotoMarkdown sends a photo with MarkdownV2 caption.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) AnswerPhotoMarkdown(photoID, text string) *AnswerMessage { func (ctx *MessageContext) AnswerPhotoMarkdown(photoID, text string) *AnswerMessage {
return ctx.answerPhoto(photoID, text, nil, tgapi.ParseMarkdownV2) return ctx.answerPhoto(photoID, text, nil, tgapi.ParseMarkdownV2)
} }
// AnswerPhotoKeyboard sends a photo with caption and inline keyboard (plain text). // AnswerPhotoKeyboard sends a photo with caption and inline keyboard (plain text).
func (ctx *MsgContext) AnswerPhotoKeyboard(photoID, text string, kb *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) AnswerPhotoKeyboard(photoID, text string, kb *InlineKeyboard) *AnswerMessage {
return ctx.answerPhoto(photoID, text, kb, tgapi.ParseNone) return ctx.answerPhoto(photoID, text, kb, tgapi.ParseNone)
} }
// AnswerPhotoKeyboardMarkdown sends a photo with caption and inline keyboard using MarkdownV2. // AnswerPhotoKeyboardMarkdown sends a photo with caption and inline keyboard using MarkdownV2.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) AnswerPhotoKeyboardMarkdown(photoID, text string, kb *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) AnswerPhotoKeyboardMarkdown(photoID, text string, kb *InlineKeyboard) *AnswerMessage {
return ctx.answerPhoto(photoID, text, kb, tgapi.ParseMarkdownV2) return ctx.answerPhoto(photoID, text, kb, tgapi.ParseMarkdownV2)
} }
// AnswerPhotof formats a string and sends it as a photo caption (plain text). // AnswerPhotof formats a string and sends it as a photo caption (plain text).
func (ctx *MsgContext) AnswerPhotof(photoID, template string, args ...any) *AnswerMessage { func (ctx *MessageContext) AnswerPhotof(photoID, template string, args ...any) *AnswerMessage {
return ctx.answerPhoto(photoID, fmt.Sprintf(template, args...), nil, tgapi.ParseNone) return ctx.answerPhoto(photoID, fmt.Sprintf(template, args...), nil, tgapi.ParseNone)
} }
// AnswerPhotofMarkdown formats a string and sends it as a photo caption using MarkdownV2. // AnswerPhotofMarkdown formats a string and sends it as a photo caption using MarkdownV2.
// //
// ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here. // ⚠️ WARNING: User input must be escaped with tgfmt.EscapeMarkdownV2() before passing here.
func (ctx *MsgContext) AnswerPhotofMarkdown(photoID, template string, args ...any) *AnswerMessage { func (ctx *MessageContext) AnswerPhotofMarkdown(photoID, template string, args ...any) *AnswerMessage {
return ctx.answerPhoto(photoID, fmt.Sprintf(template, args...), nil, tgapi.ParseMarkdownV2) return ctx.answerPhoto(photoID, fmt.Sprintf(template, args...), nil, tgapi.ParseMarkdownV2)
} }
// Internal helper that deletes a message by ID. // Internal helper that deletes a message by ID.
func (ctx *MsgContext) delete(messageID int) { func (ctx *MessageContext) delete(messageID int) {
if messageID == 0 { if messageID == 0 {
ctx.Logger.Errorln(ErrMessageIDZero) ctx.Logger.Errorln(ErrMessageIDZero)
return return
@@ -465,7 +466,7 @@ func (ctx *MsgContext) delete(messageID int) {
func (m *AnswerMessage) Delete() { m.ctx.delete(m.MessageID) } func (m *AnswerMessage) Delete() { m.ctx.delete(m.MessageID) }
// CallbackDelete deletes the message that triggered the callback query. // CallbackDelete deletes the message that triggered the callback query.
func (ctx *MsgContext) CallbackDelete() { func (ctx *MessageContext) CallbackDelete() {
if ctx.CallbackMsgID == 0 { if ctx.CallbackMsgID == 0 {
ctx.Logger.Errorln(ErrCallbackMessageMissing) ctx.Logger.Errorln(ErrCallbackMessageMissing)
return return
@@ -474,7 +475,7 @@ func (ctx *MsgContext) CallbackDelete() {
} }
// Internal helper that answers a callback query with optional text, alert, or URL. // Internal helper that answers a callback query with optional text, alert, or URL.
func (ctx *MsgContext) answerCallbackQuery(url, text string, showAlert bool) { func (ctx *MessageContext) answerCallbackQuery(url, text string, showAlert bool) {
if len(ctx.CallbackQueryID) == 0 { if len(ctx.CallbackQueryID) == 0 {
return return
} }
@@ -488,19 +489,19 @@ func (ctx *MsgContext) answerCallbackQuery(url, text string, showAlert bool) {
} }
// AnswerCallback answers the callback query with no text or alert. // AnswerCallback answers the callback query with no text or alert.
func (ctx *MsgContext) AnswerCallback() { ctx.answerCallbackQuery("", "", false) } func (ctx *MessageContext) AnswerCallback() { ctx.answerCallbackQuery("", "", false) }
// AnswerCallbackText answers the callback query with a text notification. // AnswerCallbackText answers the callback query with a text notification.
func (ctx *MsgContext) AnswerCallbackText(text string) { ctx.answerCallbackQuery("", text, false) } func (ctx *MessageContext) AnswerCallbackText(text string) { ctx.answerCallbackQuery("", text, false) }
// AnswerCallbackAlert answers the callback query with a user-visible alert. // AnswerCallbackAlert answers the callback query with a user-visible alert.
func (ctx *MsgContext) AnswerCallbackAlert(text string) { ctx.answerCallbackQuery("", text, true) } func (ctx *MessageContext) AnswerCallbackAlert(text string) { ctx.answerCallbackQuery("", text, true) }
// AnswerCallbackURL answers the callback query with a URL redirect. // AnswerCallbackURL answers the callback query with a URL redirect.
func (ctx *MsgContext) AnswerCallbackURL(u string) { ctx.answerCallbackQuery(u, "", false) } func (ctx *MessageContext) AnswerCallbackURL(u string) { ctx.answerCallbackQuery(u, "", false) }
// SendAction sends a chat action (typing, uploading_photo, etc.) to indicate bot activity. // SendAction sends a chat action (typing, uploading_photo, etc.) to indicate bot activity.
func (ctx *MsgContext) SendAction(action tgapi.ChatActionType) { func (ctx *MessageContext) SendAction(action tgapi.ChatActionType) {
if ctx.Msg == nil { if ctx.Msg == nil {
ctx.Logger.Errorln("Can't send action without chat message context") ctx.Logger.Errorln("Can't send action without chat message context")
return return
@@ -518,12 +519,12 @@ func (ctx *MsgContext) SendAction(action tgapi.ChatActionType) {
} }
// Internal helper that formats, sends, and logs an error. // Internal helper that formats, sends, and logs an error.
func (ctx *MsgContext) error(err error) { func (ctx *MessageContext) error(err error) {
if err == nil { if err == nil {
return return
} }
ctx.Logger.Errorln(err) ctx.Logger.Errorln(err)
if IsInternalError(err) { if !IsUserError(err) {
return return
} }
text := fmt.Sprintf(ctx.errorTemplate, err.Error()) text := fmt.Sprintf(ctx.errorTemplate, err.Error())
@@ -536,9 +537,9 @@ func (ctx *MsgContext) error(err error) {
} }
// Error is an alias for error(). // Error is an alias for error().
func (ctx *MsgContext) Error(err error) { ctx.error(err) } func (ctx *MessageContext) Error(err error) { ctx.error(err) }
func (ctx *MsgContext) newDraft(parseMode tgapi.ParseMode) *Draft { func (ctx *MessageContext) newDraft(parseMode tgapi.ParseMode) *Draft {
if ctx.Msg == nil { if ctx.Msg == nil {
ctx.Logger.Errorln(ErrMessageContextNil) ctx.Logger.Errorln(ErrMessageContextNil)
return nil return nil
@@ -567,20 +568,20 @@ func (ctx *MsgContext) newDraft(parseMode tgapi.ParseMode) *Draft {
// NewDraft creates a new message draft associated with the current chat. // NewDraft creates a new message draft associated with the current chat.
// Uses the API limiter to avoid rate limiting. // Uses the API limiter to avoid rate limiting.
func (ctx *MsgContext) NewDraft() *Draft { func (ctx *MessageContext) NewDraft() *Draft {
return ctx.newDraft(tgapi.ParseNone) return ctx.newDraft(tgapi.ParseNone)
} }
// NewDraftMarkdown creates a new message draft associated with the current chat, // NewDraftMarkdown creates a new message draft associated with the current chat,
// with Markdown V2 parse mode enabled. // with Markdown V2 parse mode enabled.
// Uses the API limiter to avoid rate limiting. // Uses the API limiter to avoid rate limiting.
func (ctx *MsgContext) NewDraftMarkdown() *Draft { func (ctx *MessageContext) NewDraftMarkdown() *Draft {
return ctx.newDraft(tgapi.ParseMarkdownV2) return ctx.newDraft(tgapi.ParseMarkdownV2)
} }
// Translate looks up a key in the current user's language. // Translate looks up a key in the current user's language.
// Falls back to the bot's default language if user's language is unknown or unsupported. // Falls back to the bot's default language if user's language is unknown or unsupported.
func (ctx *MsgContext) Translate(key string) string { func (ctx *MessageContext) Translate(key string) string {
if ctx.From == nil { if ctx.From == nil {
return key return key
} }
@@ -590,12 +591,12 @@ func (ctx *MsgContext) Translate(key string) string {
// NewInlineKeyboard creates a new keyboard builder with the context's payload // NewInlineKeyboard creates a new keyboard builder with the context's payload
// encoding type and the specified maximum number of buttons per row. // encoding type and the specified maximum number of buttons per row.
func (ctx *MsgContext) NewInlineKeyboard(maxRow int) *InlineKeyboard { func (ctx *MessageContext) NewInlineKeyboard(maxRow int) *InlineKeyboard {
return NewInlineKeyboard(ctx.payloadType, maxRow) return NewInlineKeyboard(ctx.payloadType, maxRow)
} }
// NewInlineKeyboardButton creates a button builder using the context payload encoding. // NewInlineKeyboardButton creates a button builder using the context payload encoding.
func (ctx *MsgContext) NewInlineKeyboardButton(text string) InlineKeyboardButtonBuilder { func (ctx *MessageContext) NewInlineKeyboardButton(text string) InlineKeyboardButtonBuilder {
return NewInlineKeyboardButton(text).SetPayloadType(ctx.payloadType) return NewInlineKeyboardButton(text).SetPayloadType(ctx.payloadType)
} }
@@ -684,19 +685,19 @@ func bindPositional(args []string, dst any) error {
// are provided than fields, the remaining fields keep their zero values. If the // are provided than fields, the remaining fields keep their zero values. If the
// final bindable field is a string, it receives the remaining arguments joined // final bindable field is a string, it receives the remaining arguments joined
// with spaces. // with spaces.
func (ctx *MsgContext) BindArgs(dst any) error { func (ctx *MessageContext) BindArgs(dst any) error {
return bindPositional(ctx.Args, dst) return bindPositional(ctx.Args, dst)
} }
// Context returns the request-scoped context associated with the current update. // Context returns the request-scoped context associated with the current update.
func (ctx *MsgContext) Context() context.Context { func (ctx *MessageContext) Context() context.Context {
if ctx.ctx == nil { if ctx.ctx == nil {
return context.Background() return context.Background()
} }
return ctx.ctx return ctx.ctx
} }
func (ctx *MsgContext) emitPolicyChecked(event PolicyCheckedEvent) { func (ctx *MessageContext) emitPolicyChecked(event PolicyCheckedEvent) {
if ctx == nil || ctx.observer == nil { if ctx == nil || ctx.observer == nil {
return return
} }
@@ -713,7 +714,7 @@ func (ctx *MsgContext) emitPolicyChecked(event PolicyCheckedEvent) {
} }
// EnterScene enters the named scene at its configured entry step. // EnterScene enters the named scene at its configured entry step.
func (ctx *MsgContext) EnterScene(name string) error { func (ctx *MessageContext) EnterScene(name string) error {
if ctx.sceneRuntime == nil { if ctx.sceneRuntime == nil {
return ErrSceneRuntimeNil return ErrSceneRuntimeNil
} }
@@ -723,7 +724,7 @@ func (ctx *MsgContext) EnterScene(name string) error {
return ErrSceneNotFound return ErrSceneNotFound
} }
key, ok := ctx.sceneRuntime.buildSceneKey(scene.Scope, ctx) key, ok := buildSceneKey(scene.Scope, ctx)
if !ok { if !ok {
return ErrCantFindSession return ErrCantFindSession
} }
@@ -743,7 +744,7 @@ func (ctx *MsgContext) EnterScene(name string) error {
} }
// EnterSceneStep enters the named scene at a specific step. // EnterSceneStep enters the named scene at a specific step.
func (ctx *MsgContext) EnterSceneStep(name, step string) error { func (ctx *MessageContext) EnterSceneStep(name, step string) error {
if ctx.sceneRuntime == nil { if ctx.sceneRuntime == nil {
return ErrSceneRuntimeNil return ErrSceneRuntimeNil
} }
@@ -756,7 +757,7 @@ func (ctx *MsgContext) EnterSceneStep(name, step string) error {
return ErrSceneStepNotFound return ErrSceneStepNotFound
} }
key, ok := ctx.sceneRuntime.buildSceneKey(scene.Scope, ctx) key, ok := buildSceneKey(scene.Scope, ctx)
if !ok { if !ok {
return ErrCantFindSession return ErrCantFindSession
} }
@@ -767,7 +768,7 @@ func (ctx *MsgContext) EnterSceneStep(name, step string) error {
} }
// ExitScene leaves the currently active scene for this context. // ExitScene leaves the currently active scene for this context.
func (ctx *MsgContext) ExitScene() error { func (ctx *MessageContext) ExitScene() error {
if ctx.sceneRuntime == nil { if ctx.sceneRuntime == nil {
return ErrSceneRuntimeNil return ErrSceneRuntimeNil
} }
@@ -785,7 +786,7 @@ func (ctx *MsgContext) ExitScene() error {
return ErrSceneNotFound return ErrSceneNotFound
} }
key, ok := ctx.sceneRuntime.buildSceneKey(scene.Scope, ctx) key, ok := buildSceneKey(scene.Scope, ctx)
if !ok { if !ok {
return ErrCantFindSession return ErrCantFindSession
} }
@@ -794,16 +795,16 @@ func (ctx *MsgContext) ExitScene() error {
} }
// IsCallback reports whether the context belongs to a callback query. // IsCallback reports whether the context belongs to a callback query.
func (ctx *MsgContext) IsCallback() bool { func (ctx *MessageContext) IsCallback() bool {
return ctx.CallbackQueryID != "" || ctx.CallbackMsgID > 0 || ctx.InlineMsgID != "" return ctx.CallbackQueryID != "" || ctx.CallbackMsgID > 0 || ctx.InlineMsgID != ""
} }
// HasPhoto reports whether the current message contains a photo payload. // HasPhoto reports whether the current message contains a photo payload.
func (ctx *MsgContext) HasPhoto() bool { func (ctx *MessageContext) HasPhoto() bool {
return ctx.Msg != nil && ctx.Msg.Photo.Len() > 0 return ctx.Msg != nil && ctx.Msg.Photo.Len() > 0
} }
func (ctx *MsgContext) upsertKeyboard(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage { func (ctx *MessageContext) upsertKeyboard(text string, keyboard *InlineKeyboard, parseMode tgapi.ParseMode) *AnswerMessage {
if ctx.IsCallback() { if ctx.IsCallback() {
if ctx.HasPhoto() { if ctx.HasPhoto() {
ctx.CallbackDelete() ctx.CallbackDelete()
@@ -815,11 +816,11 @@ func (ctx *MsgContext) upsertKeyboard(text string, keyboard *InlineKeyboard, par
} }
// UpsertKeyboard edits a callback message or sends a new plain-text message with a keyboard. // UpsertKeyboard edits a callback message or sends a new plain-text message with a keyboard.
func (ctx *MsgContext) UpsertKeyboard(text string, keyboard *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) UpsertKeyboard(text string, keyboard *InlineKeyboard) *AnswerMessage {
return ctx.upsertKeyboard(text, keyboard, tgapi.ParseNone) return ctx.upsertKeyboard(text, keyboard, tgapi.ParseNone)
} }
// UpsertKeyboardMarkdown edits a callback message or sends a new MarkdownV2 message with a keyboard. // UpsertKeyboardMarkdown edits a callback message or sends a new MarkdownV2 message with a keyboard.
func (ctx *MsgContext) UpsertKeyboardMarkdown(text string, keyboard *InlineKeyboard) *AnswerMessage { func (ctx *MessageContext) UpsertKeyboardMarkdown(text string, keyboard *InlineKeyboard) *AnswerMessage {
return ctx.upsertKeyboard(text, keyboard, tgapi.ParseMarkdownV2) return ctx.upsertKeyboard(text, keyboard, tgapi.ParseMarkdownV2)
} }
+53 -22
View File
@@ -44,7 +44,7 @@ func TestAnswerPhotoIncludesDirectMessagesTopicID(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Msg: &tgapi.Message{ Msg: &tgapi.Message{
Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate},
@@ -56,6 +56,7 @@ func TestAnswerPhotoIncludesDirectMessagesTopicID(t *testing.T) {
answer := ctx.AnswerPhoto("photo-id", "caption") answer := ctx.AnswerPhoto("photo-id", "caption")
if answer == nil { if answer == nil {
t.Fatal("expected answer message") t.Fatal("expected answer message")
return
} }
if answer.MessageID != 9 { if answer.MessageID != 9 {
t.Fatalf("unexpected message id: %d", answer.MessageID) t.Fatalf("unexpected message id: %d", answer.MessageID)
@@ -73,7 +74,7 @@ func TestBindArgsBindsScalarFields(t *testing.T) {
Name string Name string
} }
ctx := &MsgContext{Args: []string{"42", "true", "3.5", "Ada", "Lovelace"}} ctx := &MessageContext{Args: []string{"42", "true", "3.5", "Ada", "Lovelace"}}
var got input var got input
if err := ctx.BindArgs(&got); err != nil { if err := ctx.BindArgs(&got); err != nil {
@@ -92,7 +93,7 @@ func TestBindArgsBindsScalarFields(t *testing.T) {
} }
func TestNewInlineKeyboardButtonUsesContextPayloadType(t *testing.T) { func TestNewInlineKeyboardButtonUsesContextPayloadType(t *testing.T) {
ctx := &MsgContext{payloadType: BotPayloadBase64} ctx := &MessageContext{payloadType: BotPayloadBase64}
kb := NewInlineKeyboardJSON(1). kb := NewInlineKeyboardJSON(1).
AddButton(ctx.NewInlineKeyboardButton("A").SetCallbackData("cmd", 1, "two")) AddButton(ctx.NewInlineKeyboardButton("A").SetCallbackData("cmd", 1, "two"))
@@ -115,7 +116,7 @@ func TestBindArgsLeavesTrailingFieldsZeroWhenArgsRunOut(t *testing.T) {
Admin bool Admin bool
} }
ctx := &MsgContext{Args: []string{"7"}} ctx := &MessageContext{Args: []string{"7"}}
var got input var got input
if err := ctx.BindArgs(&got); err != nil { if err := ctx.BindArgs(&got); err != nil {
@@ -134,7 +135,7 @@ func TestBindArgsLeavesTrailingFieldsZeroWhenArgsRunOut(t *testing.T) {
} }
func TestBindArgsRejectsInvalidTargets(t *testing.T) { func TestBindArgsRejectsInvalidTargets(t *testing.T) {
ctx := &MsgContext{Args: []string{"1"}} ctx := &MessageContext{Args: []string{"1"}}
if err := ctx.BindArgs(nil); !errors.Is(err, ErrBindArgsTargetNotPointer) { if err := ctx.BindArgs(nil); !errors.Is(err, ErrBindArgsTargetNotPointer) {
t.Fatalf("expected ErrBindArgsTargetNotPointer for nil target, got %v", err) t.Fatalf("expected ErrBindArgsTargetNotPointer for nil target, got %v", err)
@@ -151,7 +152,7 @@ func TestBindArgsReportsConversionFailures(t *testing.T) {
ID int ID int
} }
ctx := &MsgContext{Args: []string{"oops"}} ctx := &MessageContext{Args: []string{"oops"}}
var got input var got input
err := ctx.BindArgs(&got) err := ctx.BindArgs(&got)
@@ -171,7 +172,7 @@ func TestBindArgsRejectsUnsupportedFieldTypes(t *testing.T) {
Tags []string Tags []string
} }
ctx := &MsgContext{Args: []string{"tag"}} ctx := &MessageContext{Args: []string{"tag"}}
var got input var got input
err := ctx.BindArgs(&got) err := ctx.BindArgs(&got)
@@ -183,7 +184,37 @@ func TestBindArgsRejectsUnsupportedFieldTypes(t *testing.T) {
} }
} }
func TestErrorDefaultRemainsUserVisibleForMessageFlow(t *testing.T) { func TestErrorDefaultStaysInternalForMessageFlow(t *testing.T) {
client := &http.Client{
Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
t.Fatal("unexpected HTTP request for unclassified error")
return nil, nil
}),
}
api := tgapi.NewAPI(
tgapi.NewAPIOpts("token").
SetAPIURL("https://example.test").
SetHTTPClient(client),
)
defer func() {
if err := api.Close(); err != nil {
t.Fatalf("Close returned error: %v", err)
}
}()
ctx := &MessageContext{
API: api,
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}},
Logger: sneklog.NewLogger(),
errorTemplate: "Error: %s",
}
// Unclassified errors must not leak to the user. Only AsUserError replies.
ctx.error(errors.New("boom"))
}
func TestErrorUserVisibleAnswersForMessageFlow(t *testing.T) {
var requests int var requests int
var gotBody map[string]any var gotBody map[string]any
@@ -216,14 +247,14 @@ func TestErrorDefaultRemainsUserVisibleForMessageFlow(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}},
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
errorTemplate: "Error: %s", errorTemplate: "Error: %s",
} }
ctx.error(errors.New("boom")) ctx.error(AsUserError(errors.New("boom")))
if requests != 1 { if requests != 1 {
t.Fatalf("expected one user-facing error reply, got %d requests", requests) t.Fatalf("expected one user-facing error reply, got %d requests", requests)
@@ -252,7 +283,7 @@ func TestErrorInternalSkipsUserReplyForMessageFlow(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}},
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
@@ -281,7 +312,7 @@ func TestErrorInternalSkipsCallbackAnswer(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
errorTemplate: "%s", errorTemplate: "%s",
@@ -324,7 +355,7 @@ func TestErrorUserVisibleAnswersCallback(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
errorTemplate: "Oops: %s", errorTemplate: "Oops: %s",
@@ -344,13 +375,13 @@ func TestErrorUserVisibleAnswersCallback(t *testing.T) {
func TestIsCallbackIncludesInlineCallbackTargets(t *testing.T) { func TestIsCallbackIncludesInlineCallbackTargets(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
ctx MsgContext ctx MessageContext
want bool want bool
}{ }{
{name: "callback query id", ctx: MsgContext{CallbackQueryID: "cb-1"}, want: true}, {name: "callback query id", ctx: MessageContext{CallbackQueryID: "cb-1"}, want: true},
{name: "callback message id", ctx: MsgContext{CallbackMsgID: 12}, want: true}, {name: "callback message id", ctx: MessageContext{CallbackMsgID: 12}, want: true},
{name: "inline message id", ctx: MsgContext{InlineMsgID: "inline-1"}, want: true}, {name: "inline message id", ctx: MessageContext{InlineMsgID: "inline-1"}, want: true},
{name: "not callback", ctx: MsgContext{}, want: false}, {name: "not callback", ctx: MessageContext{}, want: false},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -397,7 +428,7 @@ func TestUpsertKeyboardEditsInlineCallback(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
InlineMsgID: "inline-1", InlineMsgID: "inline-1",
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
@@ -426,7 +457,7 @@ func TestUpsertKeyboardEditsInlineCallback(t *testing.T) {
} }
func TestAnswerRejectsEmptyMessage(t *testing.T) { func TestAnswerRejectsEmptyMessage(t *testing.T) {
ctx := &MsgContext{ ctx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}},
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
} }
@@ -455,7 +486,7 @@ func TestAnswerRejectsLongMessageWithoutSendingRequest(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}},
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
@@ -539,7 +570,7 @@ func TestAnswerLongSplitsRequestsAndAttachesKeyboardToLastChunk(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}},
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
+5 -5
View File
@@ -7,7 +7,7 @@ import (
"git.scuroneko.dev/scuroneko/laniakea/tgapi" "git.scuroneko.dev/scuroneko/laniakea/tgapi"
) )
func (bot *Bot[T]) handleMessage(update *tgapi.Update, ctx *MsgContext) bool { func (bot *Bot[T]) handleMessage(update *tgapi.Update, ctx *MessageContext) bool {
text, ok := messageText(update) text, ok := messageText(update)
if !ok { if !ok {
return false return false
@@ -22,7 +22,7 @@ func (bot *Bot[T]) handleMessage(update *tgapi.Update, ctx *MsgContext) bool {
if strings.Contains(cmd, "@") { if strings.Contains(cmd, "@") {
botUsername := bot.username botUsername := bot.username
if botUsername != "" && strings.HasSuffix(cmd, "@"+botUsername) { if botUsername != "" && strings.HasSuffix(cmd, "@"+botUsername) {
cmd = cmd[:len(cmd)-len("@"+botUsername)] // убираем @botname cmd = cmd[:len(cmd)-len("@"+botUsername)] // remove @botname
} }
} }
// Ищем команду по точному совпадению // Ищем команду по точному совпадению
@@ -30,7 +30,7 @@ func (bot *Bot[T]) handleMessage(update *tgapi.Update, ctx *MsgContext) bool {
if _, exists := plugin.commands[cmd]; exists { if _, exists := plugin.commands[cmd]; exists {
ctx.Text = args ctx.Text = args
ctx.Args = strings.Fields(args) // Убирает лишние пробелы ctx.Args = strings.Fields(args)
if plugin.logger != nil { if plugin.logger != nil {
ctx.Logger = plugin.logger ctx.Logger = plugin.logger
@@ -90,7 +90,7 @@ func (bot *Bot[T]) handleMessage(update *tgapi.Update, ctx *MsgContext) bool {
return bot.handleFallback(update, ctx) return bot.handleFallback(update, ctx)
} }
func (bot *Bot[T]) handleFallback(update *tgapi.Update, ctx *MsgContext) bool { func (bot *Bot[T]) handleFallback(update *tgapi.Update, ctx *MessageContext) bool {
text, ok := messageText(update) text, ok := messageText(update)
if !ok { if !ok {
return false return false
@@ -180,7 +180,7 @@ func messageText(update *tgapi.Update) (string, bool) {
return text, true return text, true
} }
func (bot *Bot[T]) handleCallback(update *tgapi.Update, ctx *MsgContext) bool { func (bot *Bot[T]) handleCallback(update *tgapi.Update, ctx *MessageContext) bool {
data, err := bot.decodePayload(update.CallbackQuery.Data) data, err := bot.decodePayload(update.CallbackQuery.Data)
if err != nil { if err != nil {
bot.logger.Errorln(err) bot.logger.Errorln(err)
+32 -11
View File
@@ -23,6 +23,7 @@ type Plugin[T AppData] struct {
middlewares extypes.Slice[Middleware[T]] // Shared middlewares for all commands/payloads middlewares extypes.Slice[Middleware[T]] // Shared middlewares for all commands/payloads
skipAutoCmd bool // If true, all commands in this plugin are excluded from auto-help skipAutoCmd bool // If true, all commands in this plugin are excluded from auto-help
logger *sneklog.Logger logger *sneklog.Logger
loggerOwned bool // true when the logger was created by the bot during registration; only owned loggers are closed by Close
messageFallback CommandExecutor[T] messageFallback CommandExecutor[T]
handlers map[tgapi.UpdateType]CommandExecutor[T] handlers map[tgapi.UpdateType]CommandExecutor[T]
@@ -53,6 +54,9 @@ func (p *Plugin[T]) AddCommand(command *Command[T]) *Plugin[T] {
} }
return p return p
} }
if _, exists := p.commands[command.command]; exists && p.logger != nil {
p.logger.Warnf("command '%s' is already registered in plugin '%s'; overwriting", command.command, p.name)
}
p.commands[command.command] = command p.commands[command.command] = command
return p return p
} }
@@ -74,6 +78,9 @@ func (p *Plugin[T]) AddPayload(command *Command[T]) *Plugin[T] {
} }
return p return p
} }
if _, exists := p.payloads[command.command]; exists && p.logger != nil {
p.logger.Warnf("payload '%s' is already registered in plugin '%s'; overwriting", command.command, p.name)
}
p.payloads[command.command] = command p.payloads[command.command] = command
return p return p
} }
@@ -81,7 +88,7 @@ func (p *Plugin[T]) AddPayload(command *Command[T]) *Plugin[T] {
// Payload creates and immediately adds a new payload command to the plugin. // Payload creates and immediately adds a new payload command to the plugin.
// Returns the created payload command for further configuration. // Returns the created payload command for further configuration.
func (p *Plugin[T]) Payload(command string, exec CommandExecutor[T], args ...CommandArg) *Command[T] { func (p *Plugin[T]) Payload(command string, exec CommandExecutor[T], args ...CommandArg) *Command[T] {
cmd := NewPayload(command, exec, args...) cmd := NewCommand(command, exec, args...)
p.AddPayload(cmd) p.AddPayload(cmd)
return cmd return cmd
} }
@@ -101,6 +108,9 @@ func (p *Plugin[T]) AddScene(scene *Scene[T]) *Plugin[T] {
} }
scene.PluginName = p.name scene.PluginName = p.name
scene.setPluginName(p.name) scene.setPluginName(p.name)
if _, exists := p.scenes[scene.Name]; exists && p.logger != nil {
p.logger.Warnf("scene '%s' is already registered in plugin '%s'; overwriting", scene.Name, p.name)
}
p.scenes[scene.Name] = scene p.scenes[scene.Name] = scene
return p return p
} }
@@ -209,9 +219,13 @@ func (p *Plugin[T]) SetMessageFallback(handler CommandExecutor[T]) *Plugin[T] {
// Close releases plugin-owned resources such as its logger and optional // Close releases plugin-owned resources such as its logger and optional
// OnClose callback. // OnClose callback.
//
// Only loggers created by the bot during registration are closed. A logger
// supplied via SetLogger remains the caller's responsibility — the framework
// never closes a logger it does not own.
func (p *Plugin[T]) Close() error { func (p *Plugin[T]) Close() error {
var e []error var e []error
if p.logger != nil { if p.logger != nil && p.loggerOwned {
if err := p.logger.Close(); err != nil { if err := p.logger.Close(); err != nil {
e = append(e, err) e = append(e, err)
} }
@@ -225,7 +239,7 @@ func (p *Plugin[T]) Close() error {
} }
// Internal helper that validates and executes a command handler. // Internal helper that validates and executes a command handler.
func (p *Plugin[T]) executeCmd(cmd string, ctx *MsgContext, db T) error { func (p *Plugin[T]) executeCmd(cmd string, ctx *MessageContext, db T) error {
command, exists := p.commands[cmd] command, exists := p.commands[cmd]
if !exists { if !exists {
return AsInternalError(errCommandNotFound) return AsInternalError(errCommandNotFound)
@@ -247,7 +261,7 @@ func (p *Plugin[T]) executeCmd(cmd string, ctx *MsgContext, db T) error {
} }
// Internal helper that validates and executes a payload handler. // Internal helper that validates and executes a payload handler.
func (p *Plugin[T]) executePayload(payload string, ctx *MsgContext, db T) error { func (p *Plugin[T]) executePayload(payload string, ctx *MessageContext, db T) error {
command, exists := p.payloads[payload] command, exists := p.payloads[payload]
if !exists { if !exists {
return AsInternalError(errPayloadNotFound) return AsInternalError(errPayloadNotFound)
@@ -269,7 +283,7 @@ func (p *Plugin[T]) executePayload(payload string, ctx *MsgContext, db T) error
} }
// Internal helper that runs plugin middlewares in order. // Internal helper that runs plugin middlewares in order.
func (p *Plugin[T]) executeMiddlewares(ctx *MsgContext, db T) bool { func (p *Plugin[T]) executeMiddlewares(ctx *MessageContext, db T) bool {
for _, m := range p.middlewares { for _, m := range p.middlewares {
if !m.Execute(ctx, db) { if !m.Execute(ctx, db) {
return false return false
@@ -281,14 +295,14 @@ func (p *Plugin[T]) executeMiddlewares(ctx *MsgContext, db T) bool {
// MiddlewareExecutor is the function type for middleware logic. // MiddlewareExecutor is the function type for middleware logic.
// Returns true to continue execution, false to block it. // Returns true to continue execution, false to block it.
// If async, return value is ignored. // If async, return value is ignored.
type MiddlewareExecutor[T AppData] func(ctx *MsgContext, db T) bool type MiddlewareExecutor[T AppData] func(ctx *MessageContext, db T) bool
// Middleware represents a reusable execution interceptor. // Middleware represents a reusable execution interceptor.
// Can be synchronous (blocking) or asynchronous (non-blocking). // Can be synchronous (blocking) or asynchronous (non-blocking).
type Middleware[T AppData] struct { type Middleware[T AppData] struct {
name string // Human-readable name for logging/debugging name string // Human-readable name for logging/debugging
executor MiddlewareExecutor[T] // Function to execute executor MiddlewareExecutor[T] // Function to execute
order int // Optional sort order (not used yet) order int // Sort order for bot-level middleware ordering
async bool // If true, runs in goroutine and doesn't block async bool // If true, runs in goroutine and doesn't block
} }
@@ -313,12 +327,19 @@ func (m Middleware[T]) SetAsync(async bool) Middleware[T] {
// Execute runs the middleware. // Execute runs the middleware.
// If async, runs in a goroutine and returns true immediately. // If async, runs in a goroutine and returns true immediately.
// Otherwise, returns the result of the executor. // Otherwise, returns the result of the executor.
func (m Middleware[T]) Execute(ctx *MsgContext, db T) bool { //
// Async note: the goroutine receives a shallow copy of MessageContext, so
// scalar fields (FromID, ChatID, CallbackQueryID, ...) remain a stable
// snapshot. Pointer and slice fields (Msg, From, Chat, API, Logger, Args)
// 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.
func (m Middleware[T]) Execute(ctx *MessageContext, db T) bool {
if m.async { if m.async {
ctx := *ctx // copy context to avoid race condition ctxCopy := *ctx
go func(ctx MsgContext) { go func(ctx MessageContext) {
m.executor(&ctx, db) m.executor(&ctx, db)
}(ctx) }(ctxCopy)
return true return true
} }
return m.executor(ctx, db) return m.executor(ctx, db)
+10 -10
View File
@@ -6,7 +6,7 @@ import (
) )
func TestValidateArgsRequiresFullMatch(t *testing.T) { func TestValidateArgsRequiresFullMatch(t *testing.T) {
intCmd := NewCommand("int", func(ctx *MsgContext, db NoData) error { return nil }, NewCommandArg("n").SetValueType(CommandValueInt).SetRequired()) intCmd := NewCommand("int", func(ctx *MessageContext, db NoData) error { return nil }, NewCommandArg("n").SetValueType(CommandValueInt).SetRequired())
if err := intCmd.validateArgs([]string{"123"}); err != nil { if err := intCmd.validateArgs([]string{"123"}); err != nil {
t.Fatalf("expected valid integer argument, got %v", err) t.Fatalf("expected valid integer argument, got %v", err)
} }
@@ -14,7 +14,7 @@ func TestValidateArgsRequiresFullMatch(t *testing.T) {
t.Fatalf("expected ErrCmdArgRegexpMismatch for partial int match, got %v", err) t.Fatalf("expected ErrCmdArgRegexpMismatch for partial int match, got %v", err)
} }
boolCmd := NewCommand("bool", func(ctx *MsgContext, db NoData) error { return nil }, NewCommandArg("flag").SetValueType(CommandValueBool).SetRequired()) boolCmd := NewCommand("bool", func(ctx *MessageContext, db NoData) error { return nil }, NewCommandArg("flag").SetValueType(CommandValueBool).SetRequired())
if err := boolCmd.validateArgs([]string{"false"}); err != nil { if err := boolCmd.validateArgs([]string{"false"}); err != nil {
t.Fatalf("expected valid bool argument, got %v", err) t.Fatalf("expected valid bool argument, got %v", err)
} }
@@ -26,7 +26,7 @@ func TestValidateArgsRequiresFullMatch(t *testing.T) {
func TestValidateArgsEnforcesRequiredArgIndex(t *testing.T) { func TestValidateArgsEnforcesRequiredArgIndex(t *testing.T) {
cmd := NewCommand( cmd := NewCommand(
"mixed", "mixed",
func(ctx *MsgContext, db NoData) error { return nil }, func(ctx *MessageContext, db NoData) error { return nil },
NewCommandArg("optional"), NewCommandArg("optional"),
NewCommandArg("required").SetRequired(), NewCommandArg("required").SetRequired(),
) )
@@ -40,9 +40,9 @@ func TestValidateArgsEnforcesRequiredArgIndex(t *testing.T) {
} }
func TestCommandGroupBuildsPrefixedCommandsWithoutMutatingOriginal(t *testing.T) { func TestCommandGroupBuildsPrefixedCommandsWithoutMutatingOriginal(t *testing.T) {
groupMiddleware := NewMiddleware("group", func(ctx *MsgContext, db NoData) bool { return true }) groupMiddleware := NewMiddleware("group", func(ctx *MessageContext, db NoData) bool { return true })
commandMiddleware := NewMiddleware("command", func(ctx *MsgContext, db NoData) bool { return true }) commandMiddleware := NewMiddleware("command", func(ctx *MessageContext, db NoData) bool { return true })
cmd := NewCommand("ban", func(ctx *MsgContext, db NoData) error { return nil }). cmd := NewCommand("ban", func(ctx *MessageContext, db NoData) error { return nil }).
SetDescription("Ban user"). SetDescription("Ban user").
Use(commandMiddleware) Use(commandMiddleware)
@@ -78,9 +78,9 @@ func TestCommandGroupBuildsPrefixedCommandsWithoutMutatingOriginal(t *testing.T)
func TestCommandGroupBuildIsRepeatable(t *testing.T) { func TestCommandGroupBuildIsRepeatable(t *testing.T) {
group := NewCommandGroup[NoData]("admin"). group := NewCommandGroup[NoData]("admin").
Use(NewMiddleware("group", func(ctx *MsgContext, db NoData) bool { return true })). Use(NewMiddleware("group", func(ctx *MessageContext, db NoData) bool { return true })).
AddCommand(NewCommand("ban", func(ctx *MsgContext, db NoData) error { return nil }). AddCommand(NewCommand("ban", func(ctx *MessageContext, db NoData) error { return nil }).
Use(NewMiddleware("command", func(ctx *MsgContext, db NoData) bool { return true }))) Use(NewMiddleware("command", func(ctx *MessageContext, db NoData) bool { return true })))
first := group.Build() first := group.Build()
second := group.Build() second := group.Build()
@@ -103,7 +103,7 @@ func TestPluginCommandGroupRegistersBuiltCommands(t *testing.T) {
plugin := NewPlugin[NoData]("admin") plugin := NewPlugin[NoData]("admin")
plugin.CommandGroup("admin_", func(group *CommandGroup[NoData]) { plugin.CommandGroup("admin_", func(group *CommandGroup[NoData]) {
group.AddCommand(NewCommand("ban", func(ctx *MsgContext, db NoData) error { return nil })) group.AddCommand(NewCommand("ban", func(ctx *MessageContext, db NoData) error { return nil }))
}) })
if _, ok := plugin.commands["admin_ban"]; !ok { if _, ok := plugin.commands["admin_ban"]; !ok {
+15 -18
View File
@@ -8,11 +8,11 @@ import (
) )
// Policy defines a reusable authorization rule for the current update context. // Policy defines a reusable authorization rule for the current update context.
type Policy[T AppData] func(ctx *MsgContext, data T) error type Policy[T AppData] func(ctx *MessageContext, data T) error
// RequirePolicy adapts a Policy into a blocking middleware. // RequirePolicy adapts a Policy into a blocking middleware.
func RequirePolicy[T AppData](name string, p Policy[T]) Middleware[T] { func RequirePolicy[T AppData](name string, p Policy[T]) Middleware[T] {
return NewMiddleware(name, func(ctx *MsgContext, data T) bool { return NewMiddleware(name, func(ctx *MessageContext, data T) bool {
if err := p(ctx, data); err != nil { if err := p(ctx, data); err != nil {
ctx.emitPolicyChecked(PolicyCheckedEvent{ ctx.emitPolicyChecked(PolicyCheckedEvent{
Name: name, Name: name,
@@ -37,7 +37,7 @@ func RequirePolicy[T AppData](name string, p Policy[T]) Middleware[T] {
// AllPolicies composes policies that all must succeed. // AllPolicies composes policies that all must succeed.
func AllPolicies[T AppData](policies ...Policy[T]) Policy[T] { func AllPolicies[T AppData](policies ...Policy[T]) Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
for _, p := range policies { for _, p := range policies {
if err := p(ctx, data); err != nil { if err := p(ctx, data); err != nil {
return err return err
@@ -49,7 +49,7 @@ func AllPolicies[T AppData](policies ...Policy[T]) Policy[T] {
// AnyPolicy composes policies where at least one must succeed. // AnyPolicy composes policies where at least one must succeed.
func AnyPolicy[T AppData](policies ...Policy[T]) Policy[T] { func AnyPolicy[T AppData](policies ...Policy[T]) Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
var firstDeny error var firstDeny error
var internalErr error var internalErr error
for _, p := range policies { for _, p := range policies {
@@ -79,7 +79,7 @@ func AnyPolicy[T AppData](policies ...Policy[T]) Policy[T] {
// NotPolicy inverts a policy deny result while preserving internal failures. // NotPolicy inverts a policy deny result while preserving internal failures.
func NotPolicy[T AppData](policy Policy[T]) Policy[T] { func NotPolicy[T AppData](policy Policy[T]) Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
var err error var err error
if err = policy(ctx, data); err == nil { if err = policy(ctx, data); err == nil {
return AsUserError(errors.New("the action is not allowed due to policy violation")) return AsUserError(errors.New("the action is not allowed due to policy violation"))
@@ -93,7 +93,7 @@ func NotPolicy[T AppData](policy Policy[T]) Policy[T] {
// RequirePrivateChat allows execution only in private chats. // RequirePrivateChat allows execution only in private chats.
func RequirePrivateChat[T AppData]() Policy[T] { func RequirePrivateChat[T AppData]() Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
if ctx.Msg == nil || ctx.Msg.Chat == nil { if ctx.Msg == nil || ctx.Msg.Chat == nil {
return AsInternalError(errors.New("private-chat policy requires message chat context")) return AsInternalError(errors.New("private-chat policy requires message chat context"))
} }
@@ -108,7 +108,7 @@ func RequirePrivateChat[T AppData]() Policy[T] {
// RequireGroupChat allows execution only in group or supergroup chats. // RequireGroupChat allows execution only in group or supergroup chats.
func RequireGroupChat[T AppData]() Policy[T] { func RequireGroupChat[T AppData]() Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
if ctx.Msg == nil || ctx.Msg.Chat == nil { if ctx.Msg == nil || ctx.Msg.Chat == nil {
return AsInternalError(errors.New("group-chat policy requires message chat context")) return AsInternalError(errors.New("group-chat policy requires message chat context"))
} }
@@ -123,7 +123,7 @@ func RequireGroupChat[T AppData]() Policy[T] {
// RequireSupergroupChat allows execution only in supergroup chats. // RequireSupergroupChat allows execution only in supergroup chats.
func RequireSupergroupChat[T AppData]() Policy[T] { func RequireSupergroupChat[T AppData]() Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
if ctx.Msg == nil || ctx.Msg.Chat == nil { if ctx.Msg == nil || ctx.Msg.Chat == nil {
return AsInternalError(errors.New("supergroup-chat policy requires message chat context")) return AsInternalError(errors.New("supergroup-chat policy requires message chat context"))
} }
@@ -138,7 +138,7 @@ func RequireSupergroupChat[T AppData]() Policy[T] {
// RequireChatAdmin allows execution only for chat administrators or owners. // RequireChatAdmin allows execution only for chat administrators or owners.
func RequireChatAdmin[T AppData]() Policy[T] { func RequireChatAdmin[T AppData]() Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
if ctx.FromID == 0 || ctx.ChatID == 0 { if ctx.FromID == 0 || ctx.ChatID == 0 {
return AsInternalError(errors.New("chat-admin policy requires message chat context")) return AsInternalError(errors.New("chat-admin policy requires message chat context"))
} }
@@ -161,7 +161,7 @@ func RequireChatAdmin[T AppData]() Policy[T] {
// RequireChatCreator allows execution only for the chat owner. // RequireChatCreator allows execution only for the chat owner.
func RequireChatCreator[T AppData]() Policy[T] { func RequireChatCreator[T AppData]() Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
if ctx.FromID == 0 || ctx.ChatID == 0 { if ctx.FromID == 0 || ctx.ChatID == 0 {
return AsInternalError(errors.New("chat-creator policy requires message chat context")) return AsInternalError(errors.New("chat-creator policy requires message chat context"))
} }
@@ -184,19 +184,16 @@ func RequireChatCreator[T AppData]() Policy[T] {
// RequireBotAdmin allows execution only when the bot is an admin in the chat. // RequireBotAdmin allows execution only when the bot is an admin in the chat.
func RequireBotAdmin[T AppData]() Policy[T] { func RequireBotAdmin[T AppData]() Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
if ctx.ChatID == 0 { if ctx.ChatID == 0 {
return AsInternalError(errors.New("bot-admin policy requires message chat context")) return AsInternalError(errors.New("bot-admin policy requires message chat context"))
} }
if ctx.botID == 0 {
bot, err := ctx.API.GetMe() return AsInternalError(errors.New("bot ID is not set in context"))
if err != nil {
return AsInternalError(fmt.Errorf("failed to fetch bot info: %w", err))
} }
member, err := ctx.API.GetChatMember(tgapi.GetChatMember{ member, err := ctx.API.GetChatMember(tgapi.GetChatMember{
ChatID: ctx.ChatID, ChatID: ctx.ChatID, UserID: ctx.botID,
UserID: bot.ID,
}) })
if err != nil { if err != nil {
return AsInternalError(fmt.Errorf("failed to fetch bot member status: %w", err)) return AsInternalError(fmt.Errorf("failed to fetch bot member status: %w", err))
@@ -212,7 +209,7 @@ func RequireBotAdmin[T AppData]() Policy[T] {
// RequireCallbackFromUser allows execution only for callback queries sent by non-bot users. // RequireCallbackFromUser allows execution only for callback queries sent by non-bot users.
func RequireCallbackFromUser[T AppData]() Policy[T] { func RequireCallbackFromUser[T AppData]() Policy[T] {
return func(ctx *MsgContext, data T) error { return func(ctx *MessageContext, data T) error {
if ctx.Update.CallbackQuery == nil { if ctx.Update.CallbackQuery == nil {
return AsInternalError(errors.New("callback-user policy requires callback query context")) return AsInternalError(errors.New("callback-user policy requires callback query context"))
} }
+26 -26
View File
@@ -46,14 +46,14 @@ func TestRequirePolicyStopsExecutionOnDeniedPolicy(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}},
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
errorTemplate: "Error: %s", errorTemplate: "Error: %s",
} }
mw := RequirePolicy("deny", func(ctx *MsgContext, data NoData) error { mw := RequirePolicy("deny", func(ctx *MessageContext, data NoData) error {
return AsUserError(errors.New("blocked")) return AsUserError(errors.New("blocked"))
}) })
@@ -69,7 +69,7 @@ func TestRequirePolicyStopsExecutionOnDeniedPolicy(t *testing.T) {
} }
func TestRequirePrivateChatAllowsPrivateChat(t *testing.T) { func TestRequirePrivateChatAllowsPrivateChat(t *testing.T) {
ctx := &MsgContext{ ctx := &MessageContext{
Msg: &tgapi.Message{ Msg: &tgapi.Message{
Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate}, Chat: &tgapi.Chat{ID: 42, Type: tgapi.ChatTypePrivate},
}, },
@@ -82,7 +82,7 @@ func TestRequirePrivateChatAllowsPrivateChat(t *testing.T) {
} }
func TestRequirePrivateChatDeniesNonPrivateChat(t *testing.T) { func TestRequirePrivateChatDeniesNonPrivateChat(t *testing.T) {
ctx := &MsgContext{ ctx := &MessageContext{
Msg: &tgapi.Message{ Msg: &tgapi.Message{
Chat: &tgapi.Chat{ID: -100, Type: tgapi.ChatTypeSupergroup}, Chat: &tgapi.Chat{ID: -100, Type: tgapi.ChatTypeSupergroup},
}, },
@@ -136,7 +136,7 @@ func TestRequireChatAdminUsesNormalizedIDs(t *testing.T) {
} }
}() }()
ctx := &MsgContext{ ctx := &MessageContext{
API: api, API: api,
ChatID: -2001, ChatID: -2001,
FromID: 55, FromID: 55,
@@ -160,15 +160,15 @@ func TestRequireChatAdminUsesNormalizedIDs(t *testing.T) {
func TestAllPoliciesReturnsFirstError(t *testing.T) { func TestAllPoliciesReturnsFirstError(t *testing.T) {
want := AsUserError(errors.New("blocked")) want := AsUserError(errors.New("blocked"))
policy := AllPolicies( policy := AllPolicies(
func(ctx *MsgContext, data NoData) error { return nil }, func(ctx *MessageContext, data NoData) error { return nil },
func(ctx *MsgContext, data NoData) error { return want }, func(ctx *MessageContext, data NoData) error { return want },
func(ctx *MsgContext, data NoData) error { func(ctx *MessageContext, data NoData) error {
t.Fatal("unexpected evaluation after first failure") t.Fatal("unexpected evaluation after first failure")
return nil return nil
}, },
) )
err := policy(&MsgContext{Logger: sneklog.NewLogger()}, NoData{}) err := policy(&MessageContext{Logger: sneklog.NewLogger()}, NoData{})
if !errors.Is(err, want) { if !errors.Is(err, want) {
t.Fatalf("expected first policy error, got %v", err) t.Fatalf("expected first policy error, got %v", err)
} }
@@ -176,11 +176,11 @@ func TestAllPoliciesReturnsFirstError(t *testing.T) {
func TestAnyPolicyAllowsLaterSuccessAfterInternalError(t *testing.T) { func TestAnyPolicyAllowsLaterSuccessAfterInternalError(t *testing.T) {
policy := AnyPolicy( policy := AnyPolicy(
func(ctx *MsgContext, data NoData) error { return AsInternalError(errors.New("temporary")) }, func(ctx *MessageContext, data NoData) error { return AsInternalError(errors.New("temporary")) },
func(ctx *MsgContext, data NoData) error { return nil }, func(ctx *MessageContext, data NoData) error { return nil },
) )
if err := policy(&MsgContext{Logger: sneklog.NewLogger()}, NoData{}); err != nil { if err := policy(&MessageContext{Logger: sneklog.NewLogger()}, NoData{}); err != nil {
t.Fatalf("expected later success to allow access, got %v", err) t.Fatalf("expected later success to allow access, got %v", err)
} }
} }
@@ -188,11 +188,11 @@ func TestAnyPolicyAllowsLaterSuccessAfterInternalError(t *testing.T) {
func TestAnyPolicyReturnsInternalErrorWhenNonePass(t *testing.T) { func TestAnyPolicyReturnsInternalErrorWhenNonePass(t *testing.T) {
internal := AsInternalError(errors.New("temporary")) internal := AsInternalError(errors.New("temporary"))
policy := AnyPolicy( policy := AnyPolicy(
func(ctx *MsgContext, data NoData) error { return AsUserError(errors.New("denied")) }, func(ctx *MessageContext, data NoData) error { return AsUserError(errors.New("denied")) },
func(ctx *MsgContext, data NoData) error { return internal }, func(ctx *MessageContext, data NoData) error { return internal },
) )
err := policy(&MsgContext{Logger: sneklog.NewLogger()}, NoData{}) err := policy(&MessageContext{Logger: sneklog.NewLogger()}, NoData{})
if !errors.Is(err, internal) { if !errors.Is(err, internal) {
t.Fatalf("expected internal error, got %v", err) t.Fatalf("expected internal error, got %v", err)
} }
@@ -201,29 +201,29 @@ func TestAnyPolicyReturnsInternalErrorWhenNonePass(t *testing.T) {
func TestAnyPolicyReturnsFirstDenyWhenNoPolicyPasses(t *testing.T) { func TestAnyPolicyReturnsFirstDenyWhenNoPolicyPasses(t *testing.T) {
first := AsUserError(errors.New("first deny")) first := AsUserError(errors.New("first deny"))
policy := AnyPolicy( policy := AnyPolicy(
func(ctx *MsgContext, data NoData) error { return first }, func(ctx *MessageContext, data NoData) error { return first },
func(ctx *MsgContext, data NoData) error { return AsUserError(errors.New("second deny")) }, func(ctx *MessageContext, data NoData) error { return AsUserError(errors.New("second deny")) },
) )
err := policy(&MsgContext{Logger: sneklog.NewLogger()}, NoData{}) err := policy(&MessageContext{Logger: sneklog.NewLogger()}, NoData{})
if !errors.Is(err, first) { if !errors.Is(err, first) {
t.Fatalf("expected first deny error, got %v", err) t.Fatalf("expected first deny error, got %v", err)
} }
} }
func TestNotPolicyInvertsUserDenyButPreservesInternalErrors(t *testing.T) { func TestNotPolicyInvertsUserDenyButPreservesInternalErrors(t *testing.T) {
inverted := NotPolicy(func(ctx *MsgContext, data NoData) error { inverted := NotPolicy(func(ctx *MessageContext, data NoData) error {
return AsUserError(errors.New("denied")) return AsUserError(errors.New("denied"))
}) })
if err := inverted(&MsgContext{Logger: sneklog.NewLogger()}, NoData{}); err != nil { if err := inverted(&MessageContext{Logger: sneklog.NewLogger()}, NoData{}); err != nil {
t.Fatalf("expected inverted deny to succeed, got %v", err) t.Fatalf("expected inverted deny to succeed, got %v", err)
} }
internal := AsInternalError(errors.New("temporary")) internal := AsInternalError(errors.New("temporary"))
preserve := NotPolicy(func(ctx *MsgContext, data NoData) error { preserve := NotPolicy(func(ctx *MessageContext, data NoData) error {
return internal return internal
}) })
err := preserve(&MsgContext{Logger: sneklog.NewLogger()}, NoData{}) err := preserve(&MessageContext{Logger: sneklog.NewLogger()}, NoData{})
if !errors.Is(err, internal) { if !errors.Is(err, internal) {
t.Fatalf("expected internal error to be preserved, got %v", err) t.Fatalf("expected internal error to be preserved, got %v", err)
} }
@@ -232,7 +232,7 @@ func TestNotPolicyInvertsUserDenyButPreservesInternalErrors(t *testing.T) {
func TestRequirePolicyEmitsObserverEvents(t *testing.T) { func TestRequirePolicyEmitsObserverEvents(t *testing.T) {
t.Run("allow", func(t *testing.T) { t.Run("allow", func(t *testing.T) {
observer := &recordingObserver{} observer := &recordingObserver{}
ctx := &MsgContext{ ctx := &MessageContext{
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
ctx: context.Background(), ctx: context.Background(),
observer: observer, observer: observer,
@@ -240,7 +240,7 @@ func TestRequirePolicyEmitsObserverEvents(t *testing.T) {
ChatID: 20, ChatID: 20,
} }
mw := RequirePolicy("allow", func(ctx *MsgContext, data NoData) error { mw := RequirePolicy("allow", func(ctx *MessageContext, data NoData) error {
return nil return nil
}) })
@@ -257,14 +257,14 @@ func TestRequirePolicyEmitsObserverEvents(t *testing.T) {
t.Run("deny", func(t *testing.T) { t.Run("deny", func(t *testing.T) {
observer := &recordingObserver{} observer := &recordingObserver{}
ctx := &MsgContext{ ctx := &MessageContext{
Logger: sneklog.NewLogger(), Logger: sneklog.NewLogger(),
ctx: context.Background(), ctx: context.Background(),
observer: observer, observer: observer,
errorTemplate: "%s", errorTemplate: "%s",
} }
mw := RequirePolicy("deny", func(ctx *MsgContext, data NoData) error { mw := RequirePolicy("deny", func(ctx *MessageContext, data NoData) error {
return AsInternalError(errors.New("blocked")) return AsInternalError(errors.New("blocked"))
}) })
+23 -16
View File
@@ -161,23 +161,30 @@ func (bot *Bot[T]) ExecRunners(ctx context.Context) {
case <-ctx.Done(): case <-ctx.Done():
return return
case <-ticker.C: case <-ticker.C:
startedAt := time.Now() }
err := r.fn(bot) // When both ctx.Done() and ticker.C are ready at the same
bot.safeEmitEvent(ctx, RunnerFinishedEvent{ // time, Go's select picks one at random. Re-check ctx so a
Name: r.name, // late tick after cancellation does not fire one extra
Duration: time.Since(startedAt), // invocation past shutdown.
Err: err, if ctx.Err() != nil {
return
}
startedAt := time.Now()
err := r.fn(bot)
bot.safeEmitEvent(ctx, RunnerFinishedEvent{
Name: r.name,
Duration: time.Since(startedAt),
Err: err,
})
if err != nil {
bot.safeEmitEvent(ctx, ErrorEvent{
Plugin: "bot",
HandlerKind: HandlerRunnerKind,
HandlerName: r.name,
Err: err,
UserFacing: false,
}) })
if err != nil { bot.logger.Warnf("Runner %s failed: %s\n", r.name, err)
bot.safeEmitEvent(ctx, ErrorEvent{
Plugin: "bot",
HandlerKind: HandlerRunnerKind,
HandlerName: r.name,
Err: err,
UserFacing: false,
})
bot.logger.Warnf("Runner %s failed: %s\n", r.name, err)
}
} }
} }
}(runner) }(runner)
+3 -4
View File
@@ -15,7 +15,7 @@ type Scene[T any] struct {
Name string Name string
// Scope controls how active scene sessions are keyed and shared. // Scope controls how active scene sessions are keyed and shared.
Scope SceneScope Scope SceneScope
// Entry names the first step used by MsgContext.EnterScene. // Entry names the first step used by MessageContext.EnterScene.
Entry string Entry string
// PluginName stores the owning plugin name for scene resolution. // PluginName stores the owning plugin name for scene resolution.
PluginName string PluginName string
@@ -45,7 +45,7 @@ func (s *Scene[T]) SetScope(scope SceneScope) *Scene[T] {
return s return s
} }
// SetEntry sets the initial step entered by MsgContext.EnterScene. // SetEntry sets the initial step entered by MessageContext.EnterScene.
func (s *Scene[T]) SetEntry(step string) *Scene[T] { func (s *Scene[T]) SetEntry(step string) *Scene[T] {
s.Entry = step s.Entry = step
return s return s
@@ -260,8 +260,7 @@ type sceneRuntime interface {
getSession(key string) (SceneSession, error) getSession(key string) (SceneSession, error)
setSession(key string, session SceneSession) error setSession(key string, session SceneSession) error
deleteSession(key string) error deleteSession(key string) error
buildSceneKey(scope SceneScope, ctx *MsgContext) (string, bool) findSceneSession(ctx *MessageContext) (string, SceneSession, error)
findSceneSession(ctx *MsgContext) (string, SceneSession, error)
} }
type sceneMeta struct { type sceneMeta struct {
+2 -2
View File
@@ -1,8 +1,8 @@
package laniakea package laniakea
// SceneContext wraps MsgContext with scene session state for scene handlers. // SceneContext wraps MessageContext with scene session state for scene handlers.
type SceneContext struct { type SceneContext struct {
*MsgContext *MessageContext
sess SceneSession sess SceneSession
key string key string
} }
+7 -7
View File
@@ -7,10 +7,10 @@ import (
"time" "time"
) )
func (bot *Bot[T]) tryHandleScene(ctx *MsgContext) (bool, error) { func (bot *Bot[T]) tryHandleScene(ctx *MessageContext) (bool, error) {
key, session, err := bot.findSceneSession(ctx) key, session, err := bot.findSceneSession(ctx)
if err != nil { if err != nil {
if errors.Is(err, ErrCantFindSession) || errors.Is(err, ErrMessageNil) { if errors.Is(err, ErrCantFindSession) {
return false, nil return false, nil
} }
return false, err return false, err
@@ -31,9 +31,9 @@ func (bot *Bot[T]) tryHandleScene(ctx *MsgContext) (bool, error) {
return false, nil return false, nil
} }
sceneCtx := &SceneContext{ sceneCtx := &SceneContext{
MsgContext: ctx, MessageContext: ctx,
sess: session, sess: session,
key: key, key: key,
} }
return bot.executeScene(sceneCtx, scene) return bot.executeScene(sceneCtx, scene)
@@ -42,7 +42,7 @@ func (bot *Bot[T]) tryHandleScene(ctx *MsgContext) (bool, error) {
} }
func (bot *Bot[T]) executeScene(ctx *SceneContext, scene *Scene[T]) (bool, error) { func (bot *Bot[T]) executeScene(ctx *SceneContext, scene *Scene[T]) (bool, error) {
if ctx.MsgContext == nil || ctx.sess.Scene == "" { if ctx.MessageContext == nil || ctx.sess.Scene == "" {
return false, nil return false, nil
} }
@@ -269,7 +269,7 @@ func (bot *Bot[T]) applySceneResult(scene *Scene[T], ctx *SceneContext, result S
return false, nil return false, nil
} }
} }
func buildSceneKey(scope SceneScope, ctx *MsgContext) (string, bool) { func buildSceneKey(scope SceneScope, ctx *MessageContext) (string, bool) {
if ctx == nil { if ctx == nil {
return "", false return "", false
} }
+31 -31
View File
@@ -71,7 +71,7 @@ func TestBotAddPluginsPreservesScenesAndHandlesThem(t *testing.T) {
t.Fatalf("unexpected scene entry: got %q want %q", sceneMeta.Entry, "start") t.Fatalf("unexpected scene entry: got %q want %q", sceneMeta.Entry, "start")
} }
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -95,7 +95,7 @@ func TestBotAddPluginsPreservesScenesAndHandlesThem(t *testing.T) {
t.Fatal("expected scene step handler to be called") t.Fatal("expected scene step handler to be called")
} }
lookupCtx := &MsgContext{ lookupCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
} }
@@ -108,7 +108,7 @@ func TestBuildSceneKeyRejectsMissingContextFields(t *testing.T) {
tests := []struct { tests := []struct {
name string name string
scope SceneScope scope SceneScope
ctx *MsgContext ctx *MessageContext
}{ }{
{ {
name: "nil context", name: "nil context",
@@ -118,17 +118,17 @@ func TestBuildSceneKeyRejectsMissingContextFields(t *testing.T) {
{ {
name: "missing message for chat scope", name: "missing message for chat scope",
scope: SceneScopeChat, scope: SceneScopeChat,
ctx: &MsgContext{}, ctx: &MessageContext{},
}, },
{ {
name: "missing from id for user scope", name: "missing from id for user scope",
scope: SceneScopeUser, scope: SceneScopeUser,
ctx: &MsgContext{}, ctx: &MessageContext{},
}, },
{ {
name: "missing from id for user chat scope", name: "missing from id for user chat scope",
scope: SceneScopeUserChat, scope: SceneScopeUserChat,
ctx: &MsgContext{ ctx: &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
}, },
}, },
@@ -155,7 +155,7 @@ func TestEnterSceneRejectsMissingEntryConfiguration(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
ctx := &MsgContext{ ctx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -178,7 +178,7 @@ func TestEnterSceneRejectsMissingEntryConfiguration(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
ctx := &MsgContext{ ctx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -192,7 +192,7 @@ func TestEnterSceneRejectsMissingEntryConfiguration(t *testing.T) {
} }
func TestSceneContextMethodsRequireRuntime(t *testing.T) { func TestSceneContextMethodsRequireRuntime(t *testing.T) {
ctx := &MsgContext{} ctx := &MessageContext{}
if err := ctx.EnterScene("signup"); !errors.Is(err, ErrSceneRuntimeNil) { if err := ctx.EnterScene("signup"); !errors.Is(err, ErrSceneRuntimeNil) {
t.Fatalf("expected ErrSceneRuntimeNil from EnterScene, got %v", err) t.Fatalf("expected ErrSceneRuntimeNil from EnterScene, got %v", err)
@@ -238,7 +238,7 @@ func TestSceneCommandHandlerRunsBeforeStep(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -287,7 +287,7 @@ func TestSceneCommandObserverEmitsLifecycleEvents(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -339,7 +339,7 @@ func TestSceneStepObserverEmitsLifecycleEvents(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -394,7 +394,7 @@ func TestSceneMessageObserverEmitsLifecycleEvents(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
key, ok := buildSceneKey(SceneScopeUserChat, &MsgContext{ key, ok := buildSceneKey(SceneScopeUserChat, &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
}) })
@@ -460,7 +460,7 @@ func TestScenePayloadHandlerRunsBeforeStep(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -517,7 +517,7 @@ func TestScenePayloadObserverEmitsLifecycleEvents(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -563,7 +563,7 @@ func TestSceneUnmatchedPayloadFallsThroughWithoutRunningStep(t *testing.T) {
stepCalled := false stepCalled := false
plugin := NewPlugin[NoData]("wizard") plugin := NewPlugin[NoData]("wizard")
plugin.Payload("ping", func(ctx *MsgContext, db NoData) error { return nil }) plugin.Payload("ping", func(ctx *MessageContext, db NoData) error { return nil })
plugin.Scene("signup"). plugin.Scene("signup").
SetEntry("start"). SetEntry("start").
OnStep("start", func(ctx *SceneContext, db NoData) (SceneResult, error) { OnStep("start", func(ctx *SceneContext, db NoData) (SceneResult, error) {
@@ -579,7 +579,7 @@ func TestSceneUnmatchedPayloadFallsThroughWithoutRunningStep(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -588,7 +588,7 @@ func TestSceneUnmatchedPayloadFallsThroughWithoutRunningStep(t *testing.T) {
t.Fatalf("EnterScene returned error: %v", err) t.Fatalf("EnterScene returned error: %v", err)
} }
key, ok := buildSceneKey(SceneScopeUserChat, &MsgContext{ key, ok := buildSceneKey(SceneScopeUserChat, &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
}) })
@@ -632,7 +632,7 @@ func TestScenePassDoesNotPersistSessionData(t *testing.T) {
commandCalled := false commandCalled := false
plugin := NewPlugin[NoData]("wizard") plugin := NewPlugin[NoData]("wizard")
plugin.Command("ping", func(ctx *MsgContext, db NoData) error { plugin.Command("ping", func(ctx *MessageContext, db NoData) error {
commandCalled = true commandCalled = true
return nil return nil
}) })
@@ -655,7 +655,7 @@ func TestScenePassDoesNotPersistSessionData(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -664,7 +664,7 @@ func TestScenePassDoesNotPersistSessionData(t *testing.T) {
t.Fatalf("EnterScene returned error: %v", err) t.Fatalf("EnterScene returned error: %v", err)
} }
key, ok := buildSceneKey(SceneScopeUserChat, &MsgContext{ key, ok := buildSceneKey(SceneScopeUserChat, &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
}) })
@@ -718,7 +718,7 @@ func TestSceneUnmatchedCommandFallsThroughWithoutRunningStep(t *testing.T) {
stepCalled = true stepCalled = true
return ctx.Stay(), nil return ctx.Stay(), nil
}) })
plugin.Command("ping", func(ctx *MsgContext, db NoData) error { plugin.Command("ping", func(ctx *MessageContext, db NoData) error {
commandCalled = true commandCalled = true
return nil return nil
}) })
@@ -731,7 +731,7 @@ func TestSceneUnmatchedCommandFallsThroughWithoutRunningStep(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -740,7 +740,7 @@ func TestSceneUnmatchedCommandFallsThroughWithoutRunningStep(t *testing.T) {
t.Fatalf("EnterScene returned error: %v", err) t.Fatalf("EnterScene returned error: %v", err)
} }
key, ok := buildSceneKey(SceneScopeUserChat, &MsgContext{ key, ok := buildSceneKey(SceneScopeUserChat, &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
}) })
@@ -800,7 +800,7 @@ func TestSceneMessageFallbackRunsWhenNoCommandOrStepMatch(t *testing.T) {
} }
bot.AddPlugins(plugin) bot.AddPlugins(plugin)
enterCtx := &MsgContext{ enterCtx := &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
sceneRuntime: bot, sceneRuntime: bot,
@@ -809,7 +809,7 @@ func TestSceneMessageFallbackRunsWhenNoCommandOrStepMatch(t *testing.T) {
t.Fatalf("EnterScene returned error: %v", err) t.Fatalf("EnterScene returned error: %v", err)
} }
key, ok := buildSceneKey(SceneScopeUserChat, &MsgContext{ key, ok := buildSceneKey(SceneScopeUserChat, &MessageContext{
Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}}, Msg: &tgapi.Message{Chat: &tgapi.Chat{ID: 100, Type: tgapi.ChatTypePrivate}},
FromID: 42, FromID: 42,
}) })
@@ -847,7 +847,7 @@ func TestFindSceneSessionSupportsUserScopeWithoutMessage(t *testing.T) {
t.Fatalf("Set returned error: %v", err) t.Fatalf("Set returned error: %v", err)
} }
key, session, err := bot.findSceneSession(&MsgContext{FromID: 42}) key, session, err := bot.findSceneSession(&MessageContext{FromID: 42})
if err != nil { if err != nil {
t.Fatalf("findSceneSession returned error: %v", err) t.Fatalf("findSceneSession returned error: %v", err)
} }
@@ -870,7 +870,7 @@ func TestSceneStoreErrorsPropagate(t *testing.T) {
sceneScopePriority: []SceneScope{SceneScopeUser}, sceneScopePriority: []SceneScope{SceneScopeUser},
} }
_, _, err := bot.findSceneSession(&MsgContext{FromID: 42}) _, _, err := bot.findSceneSession(&MessageContext{FromID: 42})
if !errors.Is(err, getErr) { if !errors.Is(err, getErr) {
t.Fatalf("expected getErr, got %v", err) t.Fatalf("expected getErr, got %v", err)
} }
@@ -887,9 +887,9 @@ func TestSceneStoreErrorsPropagate(t *testing.T) {
} }
_, err := bot.applySceneResult(scene, &SceneContext{ _, err := bot.applySceneResult(scene, &SceneContext{
MsgContext: &MsgContext{}, MessageContext: &MessageContext{},
sess: SceneSession{Scene: "signup", Step: "start"}, sess: SceneSession{Scene: "signup", Step: "start"},
key: "user_id:42:chat_id:100", key: "user_id:42:chat_id:100",
}, SceneResult{Action: SceneActionStay}) }, SceneResult{Action: SceneActionStay})
if !errors.Is(err, setErr) { if !errors.Is(err, setErr) {
t.Fatalf("expected setErr, got %v", err) t.Fatalf("expected setErr, got %v", err)
+2 -2
View File
@@ -6,7 +6,7 @@ import (
"git.scuroneko.dev/scuroneko/laniakea/tgapi" "git.scuroneko.dev/scuroneko/laniakea/tgapi"
) )
func (bot *Bot[T]) handleUpdate(u *tgapi.Update, ctx *MsgContext) bool { func (bot *Bot[T]) handleUpdate(u *tgapi.Update, ctx *MessageContext) bool {
handled := false handled := false
for _, plugin := range bot.plugins { for _, plugin := range bot.plugins {
handler, ok := plugin.handlers[u.Type] handler, ok := plugin.handlers[u.Type]
@@ -66,7 +66,7 @@ func (bot *Bot[T]) handleUpdate(u *tgapi.Update, ctx *MsgContext) bool {
return handled return handled
} }
func (bot *Bot[T]) prepareUpdateCtx(u *tgapi.Update, ctx *MsgContext) { func (bot *Bot[T]) prepareUpdateCtx(u *tgapi.Update, ctx *MessageContext) {
var from *tgapi.User var from *tgapi.User
var chat *tgapi.Chat var chat *tgapi.Chat
switch u.Type { switch u.Type {
+50 -5
View File
@@ -16,6 +16,10 @@ var ErrDropOverflow = errors.New("drop overflow limit")
// It supports two modes: // It supports two modes:
// - "drop" mode: immediately reject if limits are exceeded. // - "drop" mode: immediately reject if limits are exceeded.
// - "wait" mode: block until capacity is available. // - "wait" mode: block until capacity is available.
//
// Per-chat limiters are created lazily and accumulate indefinitely. Call Cleanup
// periodically (e.g. from a background runner) to evict idle entries and prevent
// unbounded memory growth in bots that serve many distinct chats.
type RateLimiter struct { type RateLimiter struct {
globalLockUntil time.Time // global cooldown timestamp (set by API errors) globalLockUntil time.Time // global cooldown timestamp (set by API errors)
globalLimiter *rate.Limiter // global token bucket (30 req/sec) globalLimiter *rate.Limiter // global token bucket (30 req/sec)
@@ -23,7 +27,8 @@ type RateLimiter struct {
chatLocks map[int64]time.Time // per-chat cooldown timestamps chatLocks map[int64]time.Time // per-chat cooldown timestamps
chatLimiters map[int64]*rate.Limiter // per-chat token buckets (1 req/sec) chatLimiters map[int64]*rate.Limiter // per-chat token buckets (1 req/sec)
chatMu sync.RWMutex // protects chatLocks and chatLimiters chatLastSeen map[int64]time.Time // last access timestamp per chat, for Cleanup eviction
chatMu sync.RWMutex // protects chatLocks, chatLimiters, and chatLastSeen
} }
// NewRateLimiter creates a new RateLimiter with default limits. // NewRateLimiter creates a new RateLimiter with default limits.
@@ -34,6 +39,32 @@ func NewRateLimiter() *RateLimiter {
globalLimiter: rate.NewLimiter(30, 30), globalLimiter: rate.NewLimiter(30, 30),
chatLimiters: make(map[int64]*rate.Limiter), chatLimiters: make(map[int64]*rate.Limiter),
chatLocks: make(map[int64]time.Time), chatLocks: make(map[int64]time.Time),
chatLastSeen: make(map[int64]time.Time),
}
}
// Cleanup removes per-chat limiter state that has not been touched within
// idleThreshold and chat cooldowns whose expiry has already passed.
//
// Safe to call concurrently with Wait/Allow. Intended for periodic invocation
// from a background runner (e.g. once a minute) to bound memory in long-running
// bots that serve many distinct chats.
func (rl *RateLimiter) Cleanup(idleThreshold time.Duration) {
now := time.Now()
rl.chatMu.Lock()
defer rl.chatMu.Unlock()
for chatID, lastSeen := range rl.chatLastSeen {
if now.Sub(lastSeen) <= idleThreshold {
continue
}
delete(rl.chatLimiters, chatID)
delete(rl.chatLastSeen, chatID)
}
for chatID, until := range rl.chatLocks {
if !until.After(now) {
delete(rl.chatLocks, chatID)
}
} }
} }
@@ -228,14 +259,28 @@ func (rl *RateLimiter) waitForChatUnlock(ctx context.Context, chatID int64) erro
} }
// Internal helper that returns or creates a per-chat limiter. // Internal helper that returns or creates a per-chat limiter.
// Updates chatLastSeen so Cleanup can evict idle entries.
func (rl *RateLimiter) getChatLimiter(chatID int64) *rate.Limiter { func (rl *RateLimiter) getChatLimiter(chatID int64) *rate.Limiter {
rl.chatMu.Lock() now := time.Now()
defer rl.chatMu.Unlock()
if lim, ok := rl.chatLimiters[chatID]; ok { rl.chatMu.RLock()
lim, ok := rl.chatLimiters[chatID]
rl.chatMu.RUnlock()
if ok {
rl.chatMu.Lock()
rl.chatLastSeen[chatID] = now
rl.chatMu.Unlock()
return lim return lim
} }
lim := rate.NewLimiter(1, 1)
rl.chatMu.Lock()
defer rl.chatMu.Unlock()
if lim, ok := rl.chatLimiters[chatID]; ok {
rl.chatLastSeen[chatID] = now
return lim
}
lim = rate.NewLimiter(1, 1)
rl.chatLimiters[chatID] = lim rl.chatLimiters[chatID] = lim
rl.chatLastSeen[chatID] = now
return lim return lim
} }
+49
View File
@@ -39,3 +39,52 @@ func TestRateLimiterGlobalWaitRespectsContextCancellation(t *testing.T) {
t.Fatalf("expected DeadlineExceeded, got %v", err) t.Fatalf("expected DeadlineExceeded, got %v", err)
} }
} }
// TestRateLimiterCleanupEvictsIdleChats guards the memory-leak fix: per-chat
// limiter and lastSeen state must be reclaimed by Cleanup once the entry has
// been idle for longer than the threshold, while still-active chats and
// unexpired cooldowns must survive.
func TestRateLimiterCleanupEvictsIdleChats(t *testing.T) {
rl := NewRateLimiter()
// Touch chat 1 to make it tracked, then backdate its last-seen marker
// so it looks idle from Cleanup's perspective.
if !rl.Allow(1) {
t.Fatal("expected initial Allow for chat 1 to succeed")
}
rl.chatMu.Lock()
rl.chatLastSeen[1] = time.Now().Add(-time.Hour)
rl.chatMu.Unlock()
// Touch chat 2 so it stays "active".
if !rl.Allow(2) {
t.Fatal("expected initial Allow for chat 2 to succeed")
}
// Expired cooldown should be evicted; future cooldown should survive.
rl.SetChatLock(10, 1)
rl.chatMu.Lock()
rl.chatLocks[10] = time.Now().Add(-time.Second)
rl.chatLocks[11] = time.Now().Add(time.Hour)
rl.chatMu.Unlock()
rl.Cleanup(time.Minute)
rl.chatMu.RLock()
defer rl.chatMu.RUnlock()
if _, ok := rl.chatLimiters[1]; ok {
t.Fatal("expected idle chat 1 limiter to be evicted")
}
if _, ok := rl.chatLastSeen[1]; ok {
t.Fatal("expected idle chat 1 lastSeen to be evicted")
}
if _, ok := rl.chatLimiters[2]; !ok {
t.Fatal("expected active chat 2 limiter to remain")
}
if _, ok := rl.chatLocks[10]; ok {
t.Fatal("expected expired chat 10 lock to be evicted")
}
if _, ok := rl.chatLocks[11]; !ok {
t.Fatal("expected future chat 11 lock to remain")
}
}