fix: send forbidden reply when unauthorized user messages bot
When a user not on the allowlist sends a message, the bot now replies with "You are not authorized to use this bot." instead of silently dropping the message. Fixes the confirmed case in whatsapp_native and also covers telegram and the HandleMessage safety-net in base.go. Closes #49 Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
84e42d6904
commit
0ca27794b4
4 changed files with 90 additions and 2 deletions
|
|
@ -18,6 +18,9 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/media"
|
"github.com/sipeed/picoclaw/pkg/media"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// ForbiddenReplyText is the message sent to unauthorized users who message the bot.
|
||||||
|
const ForbiddenReplyText = "You are not authorized to use this bot."
|
||||||
|
|
||||||
var (
|
var (
|
||||||
uniqueIDCounter uint64
|
uniqueIDCounter uint64
|
||||||
uniqueIDPrefix string
|
uniqueIDPrefix string
|
||||||
|
|
@ -259,10 +262,22 @@ func (c *BaseChannel) HandleMessage(
|
||||||
}
|
}
|
||||||
if sender.CanonicalID != "" || sender.PlatformID != "" {
|
if sender.CanonicalID != "" || sender.PlatformID != "" {
|
||||||
if !c.IsAllowedSender(sender) {
|
if !c.IsAllowedSender(sender) {
|
||||||
|
logger.DebugCF(c.name, "Message blocked (not in allowlist)", map[string]any{"sender": sender.CanonicalID})
|
||||||
|
if c.owner != nil {
|
||||||
|
_, _ = c.owner.Send(ctx, bus.OutboundMessage{
|
||||||
|
Channel: c.name, ChatID: chatID, Content: ForbiddenReplyText,
|
||||||
|
})
|
||||||
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
if !c.IsAllowed(senderID) {
|
if !c.IsAllowed(senderID) {
|
||||||
|
logger.DebugCF(c.name, "Message blocked (not in allowlist)", map[string]any{"sender_id": senderID})
|
||||||
|
if c.owner != nil {
|
||||||
|
_, _ = c.owner.Send(ctx, bus.OutboundMessage{
|
||||||
|
Channel: c.name, ChatID: chatID, Content: ForbiddenReplyText,
|
||||||
|
})
|
||||||
|
}
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,12 +1,17 @@
|
||||||
package channels
|
package channels
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/bus"
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
"github.com/sipeed/picoclaw/pkg/config"
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func newTestBus() *bus.MessageBus {
|
||||||
|
return bus.NewMessageBus()
|
||||||
|
}
|
||||||
|
|
||||||
func TestBaseChannelIsAllowed(t *testing.T) {
|
func TestBaseChannelIsAllowed(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
|
|
@ -263,3 +268,65 @@ func TestIsAllowedSender(t *testing.T) {
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHandleMessageForbiddenReply(t *testing.T) {
|
||||||
|
const allowedID = "allowed:123"
|
||||||
|
const chatID = "chat-456"
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
sender bus.SenderInfo
|
||||||
|
wantReplied bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "unauthorized sender receives forbidden reply",
|
||||||
|
sender: bus.SenderInfo{
|
||||||
|
Platform: "test",
|
||||||
|
PlatformID: "999",
|
||||||
|
CanonicalID: "test:999",
|
||||||
|
},
|
||||||
|
wantReplied: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "authorized sender does not receive forbidden reply",
|
||||||
|
sender: bus.SenderInfo{
|
||||||
|
Platform: "test",
|
||||||
|
PlatformID: "123",
|
||||||
|
CanonicalID: allowedID,
|
||||||
|
},
|
||||||
|
wantReplied: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
mock := &mockChannel{}
|
||||||
|
mock.BaseChannel = *NewBaseChannel("test", nil, newTestBus(), []string{allowedID})
|
||||||
|
mock.SetOwner(mock)
|
||||||
|
|
||||||
|
mock.HandleMessage(
|
||||||
|
context.Background(),
|
||||||
|
bus.Peer{Kind: "direct", ID: chatID},
|
||||||
|
"msg-1", tt.sender.PlatformID, chatID, "hello",
|
||||||
|
nil, nil,
|
||||||
|
tt.sender,
|
||||||
|
)
|
||||||
|
|
||||||
|
if tt.wantReplied {
|
||||||
|
if len(mock.sentMessages) != 1 {
|
||||||
|
t.Fatalf("expected 1 forbidden reply, got %d", len(mock.sentMessages))
|
||||||
|
}
|
||||||
|
if mock.sentMessages[0].Content != ForbiddenReplyText {
|
||||||
|
t.Errorf("reply content = %q, want %q", mock.sentMessages[0].Content, ForbiddenReplyText)
|
||||||
|
}
|
||||||
|
if mock.sentMessages[0].ChatID != chatID {
|
||||||
|
t.Errorf("reply chatID = %q, want %q", mock.sentMessages[0].ChatID, chatID)
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
if len(mock.sentMessages) != 0 {
|
||||||
|
t.Fatalf("expected no reply for authorized sender, got %d", len(mock.sentMessages))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -592,21 +592,25 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
|
||||||
DisplayName: user.FirstName,
|
DisplayName: user.FirstName,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
chatID := message.Chat.ID
|
||||||
|
chatIDStr := fmt.Sprintf("%d", chatID)
|
||||||
|
|
||||||
// check allowlist to avoid downloading attachments for rejected users
|
// check allowlist to avoid downloading attachments for rejected users
|
||||||
if !c.IsAllowedSender(sender) {
|
if !c.IsAllowedSender(sender) {
|
||||||
logger.DebugCF("telegram", "Message rejected by allowlist", map[string]any{
|
logger.DebugCF("telegram", "Message rejected by allowlist", map[string]any{
|
||||||
"user_id": platformID,
|
"user_id": platformID,
|
||||||
})
|
})
|
||||||
|
_, _ = c.Send(ctx, bus.OutboundMessage{
|
||||||
|
Channel: "telegram", ChatID: chatIDStr, Content: channels.ForbiddenReplyText,
|
||||||
|
})
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
chatID := message.Chat.ID
|
|
||||||
c.chatIDs[platformID] = chatID
|
c.chatIDs[platformID] = chatID
|
||||||
|
|
||||||
content := ""
|
content := ""
|
||||||
mediaPaths := []string{}
|
mediaPaths := []string{}
|
||||||
|
|
||||||
chatIDStr := fmt.Sprintf("%d", chatID)
|
|
||||||
messageIDStr := fmt.Sprintf("%d", message.MessageID)
|
messageIDStr := fmt.Sprintf("%d", message.MessageID)
|
||||||
scope := channels.BuildMediaScope("telegram", chatIDStr, messageIDStr)
|
scope := channels.BuildMediaScope("telegram", chatIDStr, messageIDStr)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -385,6 +385,8 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
|
||||||
}
|
}
|
||||||
|
|
||||||
if !c.IsAllowedSender(sender) {
|
if !c.IsAllowedSender(sender) {
|
||||||
|
logger.DebugCF("whatsapp", "WhatsApp message blocked (not in allowlist)", map[string]any{"sender_id": senderID})
|
||||||
|
_, _ = c.Send(c.runCtx, bus.OutboundMessage{Channel: "whatsapp", ChatID: chatID, Content: channels.ForbiddenReplyText})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue