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:
github-actions[bot] 2026-04-10 19:19:41 +02:00
parent b99c3a06c0
commit 024282be26
4 changed files with 240 additions and 19 deletions

View file

@ -18,6 +18,9 @@ import (
"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 (
uniqueIDCounter uint64
uniqueIDPrefix string
@ -259,10 +262,22 @@ func (c *BaseChannel) HandleMessage(
}
if sender.CanonicalID != "" || sender.PlatformID != "" {
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
}
} else {
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
}
}

View file

@ -9,6 +9,10 @@ import (
"github.com/sipeed/picoclaw/pkg/config"
)
func newTestBus() *bus.MessageBus {
return bus.NewMessageBus()
}
func TestBaseChannelIsAllowed(t *testing.T) {
tests := []struct {
name string
@ -334,3 +338,65 @@ func TestObserveGroupMessage_SkipsTypingAndPlaceholder(t *testing.T) {
default:
}
}
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))
}
}
})
}
}

View file

@ -592,21 +592,25 @@ func (c *TelegramChannel) handleMessage(ctx context.Context, message *telego.Mes
DisplayName: user.FirstName,
}
chatID := message.Chat.ID
chatIDStr := fmt.Sprintf("%d", chatID)
// check allowlist to avoid downloading attachments for rejected users
if !c.IsAllowedSender(sender) {
logger.DebugCF("telegram", "Message rejected by allowlist", map[string]any{
"user_id": platformID,
})
_, _ = c.Send(ctx, bus.OutboundMessage{
Channel: "telegram", ChatID: chatIDStr, Content: channels.ForbiddenReplyText,
})
return nil
}
chatID := message.Chat.ID
c.chatIDs[platformID] = chatID
content := ""
mediaPaths := []string{}
chatIDStr := fmt.Sprintf("%d", chatID)
messageIDStr := fmt.Sprintf("%d", message.MessageID)
scope := channels.BuildMediaScope("telegram", chatIDStr, messageIDStr)

View file

@ -18,6 +18,7 @@ import (
"sync/atomic"
"time"
"github.com/google/uuid"
"github.com/mdp/qrterminal/v3"
"go.mau.fi/whatsmeow"
"go.mau.fi/whatsmeow/proto/waE2E"
@ -33,6 +34,7 @@ import (
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/identity"
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/media"
"github.com/sipeed/picoclaw/pkg/utils"
)
@ -355,11 +357,64 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
}
content = utils.SanitizeMessageContent(content)
if content == "" {
return
var mediaPaths []string
// storeMedia registers a downloaded local file with the MediaStore and returns
// a media ref (or the raw path as fallback when no store is configured).
storeMedia := func(localPath, filename string) string {
if store := c.GetMediaStore(); store != nil {
scope := channels.BuildMediaScope("whatsapp_native", chatID, evt.Info.ID)
ref, err := store.Store(localPath, media.MediaMeta{
Filename: filename,
Source: "whatsapp_native",
CleanupPolicy: media.CleanupPolicyDeleteOnCleanup,
}, scope)
if err == nil {
return ref
}
}
return localPath
}
var mediaPaths []string
if img := evt.Message.GetImageMessage(); img != nil {
if localPath := c.downloadWAMedia(c.runCtx, img, "photo.jpg"); localPath != "" {
mediaPaths = append(mediaPaths, storeMedia(localPath, "photo.jpg"))
if caption := img.GetCaption(); caption != "" {
if content != "" {
content += "\n"
}
content += caption
} else if content == "" {
content = "[image]"
}
}
}
if doc := evt.Message.GetDocumentMessage(); doc != nil {
filename := doc.GetFileName()
if filename == "" {
filename = "document"
}
if localPath := c.downloadWAMedia(c.runCtx, doc, filename); localPath != "" {
mediaPaths = append(mediaPaths, storeMedia(localPath, filename))
if content == "" {
content = "[file: " + filename + "]"
}
}
}
if audio := evt.Message.GetAudioMessage(); audio != nil {
if localPath := c.downloadWAMedia(c.runCtx, audio, "audio.ogg"); localPath != "" {
mediaPaths = append(mediaPaths, storeMedia(localPath, "audio.ogg"))
if content == "" {
content = "[voice]"
}
}
}
if content == "" && len(mediaPaths) == 0 {
return
}
metadata := make(map[string]string)
metadata["message_id"] = evt.Info.ID
@ -373,6 +428,9 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
metadata["peer_kind"] = "direct"
metadata["peer_id"] = senderID
}
if len(mediaPaths) > 0 && c.config.EchoTranscription {
metadata["echo_transcription"] = "true"
}
peerKind := "direct"
if evt.Info.Chat.Server == types.GroupServer {
@ -393,6 +451,7 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
"WhatsApp message blocked (not in allow_from)",
map[string]any{"sender_id": senderID},
)
_, _ = c.Send(c.runCtx, bus.OutboundMessage{Channel: "whatsapp", ChatID: chatID, Content: channels.ForbiddenReplyText})
return
}
@ -405,25 +464,20 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
)
if isGroup {
// Detect bot mention via ContextInfo.MentionedJID (populated for @mentions in groups).
isMentioned := false
c.mu.Lock()
botJID := c.client.Store.ID
botLID := c.client.Store.GetLID()
c.mu.Unlock()
var ctx2 *waE2E.ContextInfo
if ext := evt.Message.GetExtendedTextMessage(); ext != nil {
ctx2 = ext.GetContextInfo()
var botUsers []string
if botJID != nil && botJID.User != "" {
botUsers = append(botUsers, botJID.User)
}
if ctx2 != nil && botJID != nil {
botUser := botJID.User
for _, jid := range ctx2.GetMentionedJID() {
if strings.HasPrefix(jid, botUser+"@") || jid == botUser {
isMentioned = true
break
if botLID.User != "" {
botUsers = append(botUsers, botLID.User)
}
}
}
respond, cleaned := c.ShouldRespondInGroup(isMentioned, content)
respond, cleaned := c.ShouldRespondInGroup(isMentionedInGroup(evt.Message, content, botUsers), content)
if !respond {
c.ObserveGroupMessage(c.runCtx, peer, messageID, senderID, chatID, content, mediaPaths, metadata, sender)
return
@ -434,6 +488,48 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
c.HandleMessage(c.runCtx, peer, messageID, senderID, chatID, content, mediaPaths, metadata, sender)
}
// isMentionedInGroup returns true if the bot is @mentioned in a group message.
//
// WhatsApp LID sessions expose two identifiers (phone JID and LID); botUsers should
// contain the User portion of both so we match whichever format WhatsApp uses.
//
// Detection order:
// 1. ContextInfo.MentionedJID — the authoritative signal (ExtendedTextMessage only).
// 2. Text-based fallback — handles plain Conversation messages where ContextInfo
// is absent, which occurs in LID sessions for some WhatsApp clients.
func isMentionedInGroup(msg *waE2E.Message, content string, botUsers []string) bool {
if len(botUsers) == 0 {
return false
}
// Primary: ContextInfo.MentionedJID
if ext := msg.GetExtendedTextMessage(); ext != nil {
if ctx2 := ext.GetContextInfo(); ctx2 != nil {
for _, jid := range ctx2.GetMentionedJID() {
for _, u := range botUsers {
if strings.HasPrefix(jid, u+"@") || jid == u {
return true
}
}
}
}
}
// Fallback: @<user> in message text
for _, u := range botUsers {
if strings.Contains(content, "@"+u) {
return true
}
}
return false
}
// VoiceCapabilities reports that this channel supports ASR (speech-to-text).
func (c *WhatsAppNativeChannel) VoiceCapabilities() channels.VoiceCapabilities {
return channels.VoiceCapabilities{ASR: true, TTS: false}
}
func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessage) ([]string, error) {
if !c.IsRunning() {
return nil, channels.ErrNotRunning
@ -473,6 +569,46 @@ func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessag
return nil, nil
}
// downloadWAMedia downloads and decrypts a WhatsApp media message using the
// whatsmeow client, saves it to a temp file, and returns the local path.
// Returns "" on any error (errors are logged at WARN level).
func (c *WhatsAppNativeChannel) downloadWAMedia(
ctx context.Context,
msg whatsmeow.DownloadableMessage,
filename string,
) string {
c.mu.Lock()
client := c.client
c.mu.Unlock()
if client == nil {
return ""
}
data, err := client.Download(ctx, msg)
if err != nil {
logger.WarnCF(
"whatsapp",
"Failed to download media",
map[string]any{"error": err.Error(), "filename": filename},
)
return ""
}
mediaDir := media.TempDir()
if err := os.MkdirAll(mediaDir, 0o700); err != nil {
logger.WarnCF("whatsapp", "Failed to create media dir", map[string]any{"error": err.Error()})
return ""
}
localPath := filepath.Join(mediaDir, uuid.New().String()[:8]+"_"+utils.SanitizeFilename(filename))
if err := os.WriteFile(localPath, data, 0o600); err != nil {
logger.WarnCF("whatsapp", "Failed to write media file", map[string]any{"error": err.Error()})
return ""
}
return localPath
}
// parseJID converts a chat ID (phone number or JID string) to types.JID.
func parseJID(s string) (types.JID, error) {
s = strings.TrimSpace(s)