Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions events/events.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,4 +58,6 @@ const (
OnPurchasedPaidMedia = "onPurchasedPaidMedia"
// OnManagedBot is emitted when a managed bot update is received.
OnManagedBot = "onManagedBot"
// OnSubscription is emitted when a bot subscription update is received.
OnSubscription = "onSubscription"
)
5 changes: 5 additions & 0 deletions events/types.go
Original file line number Diff line number Diff line change
Expand Up @@ -113,3 +113,8 @@ type PurchasedPaidMediaEvent struct {
type ManagedBotEvent struct {
ManagedBot *client.ManagedBotUpdated
}

// SubscriptionEvent is emitted when a bot subscription update is received.
type SubscriptionEvent struct {
Subscription *client.BotSubscriptionUpdated
}
3 changes: 3 additions & 0 deletions handlers/handler.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,3 +75,6 @@ type PurchasedPaidMediaHandler func(ctx context.Context, event *events.Purchased

// ManagedBotHandler is a function that handles a managed bot event.
type ManagedBotHandler func(ctx context.Context, event *events.ManagedBotEvent) error

// SubscriptionHandler is a function that handles a bot subscription update event.
type SubscriptionHandler func(ctx context.Context, event *events.SubscriptionEvent) error
5 changes: 5 additions & 0 deletions handlers/registry.go
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,11 @@ func (r *Registry) OnManagedBot(handler ManagedBotHandler) eventemitter.Unsubscr
return onEvent(r, events.OnManagedBot, "OnManagedBot", handler)
}

// OnSubscription registers a handler for bot subscription update events.
func (r *Registry) OnSubscription(handler SubscriptionHandler) eventemitter.UnsubscribeFunc {
return onEvent(r, events.OnSubscription, "OnSubscription", handler)
}

func (r *Registry) onMessageEvent(
event string,
name string,
Expand Down
31 changes: 31 additions & 0 deletions handlers/registry_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -220,6 +220,37 @@ func TestRegistry(t *testing.T) {
}
})

t.Run("OnSubscription", func(t *testing.T) {
var called bool
var payload *events.SubscriptionEvent
unsub := reg.OnSubscription(func(_ context.Context, event *events.SubscriptionEvent) error {
called = true
payload = event

return nil
})

expectedPayload := &events.SubscriptionEvent{}
ee.Emit(context.Background(), events.OnSubscription, expectedPayload)
if !called {
t.Fatal("handler was not called")
}
if payload != expectedPayload {
t.Fatalf("payload mismatch: got %p, want %p", payload, expectedPayload)
}

called = false
payload = nil
unsub()
ee.Emit(context.Background(), events.OnSubscription, expectedPayload)
if called {
t.Fatal("handler was called after unsubscribe")
}
if payload != nil {
t.Fatalf("payload=%v, want nil after unsubscribe", payload)
}
})

t.Run("OnCommandName", func(t *testing.T) {
var called bool
unsub := reg.OnCommandName("start", func(_ context.Context, event *events.CommandEvent) error {
Expand Down
91 changes: 68 additions & 23 deletions listeners/classifier.go
Original file line number Diff line number Diff line change
Expand Up @@ -66,11 +66,19 @@ func classifyMessages(ctx context.Context, emitter eventemitter.EventEmitter, up

func classifyQueries(ctx context.Context, emitter eventemitter.EventEmitter, update *client.Update) {
if update.CallbackQuery != nil {
emitter.Emit(ctx, events.OnCallbackQuery, &events.CallbackQueryEvent{CallbackQuery: update.CallbackQuery})
emitter.Emit(
ctx,
events.OnCallbackQuery,
&events.CallbackQueryEvent{CallbackQuery: update.CallbackQuery},
)
}

if update.InlineQuery != nil {
emitter.Emit(ctx, events.OnInlineQuery, &events.InlineQueryEvent{InlineQuery: update.InlineQuery})
emitter.Emit(
ctx,
events.OnInlineQuery,
&events.InlineQueryEvent{InlineQuery: update.InlineQuery},
)
}

if update.ChosenInlineResult != nil {
Expand All @@ -80,7 +88,11 @@ func classifyQueries(ctx context.Context, emitter eventemitter.EventEmitter, upd
}

if update.ShippingQuery != nil {
emitter.Emit(ctx, events.OnShippingQuery, &events.ShippingQueryEvent{ShippingQuery: update.ShippingQuery})
emitter.Emit(
ctx,
events.OnShippingQuery,
&events.ShippingQueryEvent{ShippingQuery: update.ShippingQuery},
)
}

if update.PreCheckoutQuery != nil {
Expand Down Expand Up @@ -137,26 +149,42 @@ func classifyChatUpdates(ctx context.Context, emitter eventemitter.EventEmitter,
}

func classifyBusinessUpdates(ctx context.Context, emitter eventemitter.EventEmitter, update *client.Update) {
if update.BusinessConnection != nil {
emitter.Emit(ctx, events.OnBusinessConnection, &events.BusinessConnectionEvent{
BusinessConnection: update.BusinessConnection,
})
}

if update.DeletedBusinessMessages != nil {
emitter.Emit(ctx, events.OnDeletedBusinessMessages, &events.DeletedBusinessMessagesEvent{
DeletedBusinessMessages: update.DeletedBusinessMessages,
})
}

if update.PurchasedPaidMedia != nil {
emitter.Emit(ctx, events.OnPurchasedPaidMedia, &events.PurchasedPaidMediaEvent{
PurchasedPaidMedia: update.PurchasedPaidMedia,
})
}

if update.ManagedBot != nil {
emitter.Emit(ctx, events.OnManagedBot, &events.ManagedBotEvent{ManagedBot: update.ManagedBot})
for _, candidate := range []struct {
event string
payload any
}{
{
event: events.OnBusinessConnection,
payload: &events.BusinessConnectionEvent{
BusinessConnection: update.BusinessConnection,
},
},
{
event: events.OnDeletedBusinessMessages,
payload: &events.DeletedBusinessMessagesEvent{
DeletedBusinessMessages: update.DeletedBusinessMessages,
},
},
{
event: events.OnPurchasedPaidMedia,
payload: &events.PurchasedPaidMediaEvent{
PurchasedPaidMedia: update.PurchasedPaidMedia,
},
},
{
event: events.OnManagedBot,
payload: &events.ManagedBotEvent{ManagedBot: update.ManagedBot},
},
{
event: events.OnSubscription,
payload: &events.SubscriptionEvent{
Subscription: update.Subscription,
},
},
} {
if !isNilPayload(candidate.payload) {
emitter.Emit(ctx, candidate.event, candidate.payload)
}
}
}

Expand All @@ -166,3 +194,20 @@ func emitMessage(ctx context.Context, emitter eventemitter.EventEmitter, event s
Type: messagetype.Detect(message),
})
}

func isNilPayload(payload any) bool {
switch v := payload.(type) {
case *events.BusinessConnectionEvent:
return v.BusinessConnection == nil
case *events.DeletedBusinessMessagesEvent:
return v.DeletedBusinessMessages == nil
case *events.PurchasedPaidMediaEvent:
return v.PurchasedPaidMedia == nil
case *events.ManagedBotEvent:
return v.ManagedBot == nil
case *events.SubscriptionEvent:
return v.Subscription == nil
default:
return payload == nil
}
}
11 changes: 11 additions & 0 deletions listeners/classifier_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -237,6 +237,17 @@ func TestClassifier(t *testing.T) {
}
},
},
{
name: "Subscription",
update: &client.Update{Subscription: &client.BotSubscriptionUpdated{}},
event: events.OnSubscription,
assertType: func(t *testing.T, payload any) {
t.Helper()
if _, ok := payload.(*events.SubscriptionEvent); !ok {
t.Fatalf("payload type=%T, want *events.SubscriptionEvent", payload)
}
},
},
}

for _, tt := range tests {
Expand Down
Loading