feat(whatsapp_native): multimodal inbound/outbound media support

- handleIncoming: download Image/Video/Audio/Document/Sticker via whatsmeow,
  store in MediaStore, append Telegram-style annotations ([image: photo], etc.)
- Captions merged into text; text-only messages work without a connected client
- SendMedia implements channels.MediaSender (upload + Image/Video/Audio/Document)
- Tests: imageFilenameFromMime, caption preserved when download fails, media-only dropped when download fails

Made-with: Cursor
This commit is contained in:
Aditya Kalro 2026-03-20 08:55:28 -07:00
parent 938437bfc9
commit e33515faa0
2 changed files with 457 additions and 19 deletions

View file

@ -4,13 +4,18 @@ package whatsapp
import (
"context"
"database/sql"
"testing"
"time"
"go.mau.fi/whatsmeow"
"go.mau.fi/whatsmeow/proto/waE2E"
"go.mau.fi/whatsmeow/store/sqlstore"
"go.mau.fi/whatsmeow/types"
"go.mau.fi/whatsmeow/types/events"
waLog "go.mau.fi/whatsmeow/util/log"
"google.golang.org/protobuf/proto"
_ "modernc.org/sqlite"
"github.com/sipeed/picoclaw/pkg/bus"
"github.com/sipeed/picoclaw/pkg/channels"
@ -59,3 +64,126 @@ func TestHandleIncoming_DoesNotConsumeGenericCommandsLocally(t *testing.T) {
}
}
}
func TestImageFilenameFromMime(t *testing.T) {
tests := []struct {
mime string
want string
}{
{"image/png", "photo.png"},
{"image/jpeg", "photo.jpg"},
{"image/gif", "photo.gif"},
{"image/webp", "photo.webp"},
{"", "photo.jpg"},
}
for _, tt := range tests {
if got := imageFilenameFromMime(tt.mime); got != tt.want {
t.Errorf("imageFilenameFromMime(%q) = %q, want %q", tt.mime, got, tt.want)
}
}
}
// whatsappTestClient returns a disconnected whatsmeow client (valid store) for exercising handleIncoming media paths.
func whatsappTestClient(t *testing.T) *whatsmeow.Client {
t.Helper()
ctx := context.Background()
db, err := sql.Open("sqlite", "file::memory:?cache=shared&_foreign_keys=on")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
if _, err := db.ExecContext(ctx, "PRAGMA foreign_keys = ON"); err != nil {
t.Fatal(err)
}
container := sqlstore.NewWithDB(db, "sqlite", waLog.Noop)
if err := container.Upgrade(ctx); err != nil {
t.Fatal(err)
}
device, err := container.GetFirstDevice(ctx)
if err != nil {
t.Fatal(err)
}
return whatsmeow.NewClient(device, waLog.Noop)
}
func TestHandleIncoming_ImageCaptionDownloadFailsStillForwardsCaption(t *testing.T) {
messageBus := bus.NewMessageBus()
ch := &WhatsAppNativeChannel{
BaseChannel: channels.NewBaseChannel("whatsapp_native", config.WhatsAppConfig{}, messageBus, nil),
runCtx: context.Background(),
client: whatsappTestClient(t),
}
evt := &events.Message{
Info: types.MessageInfo{
MessageSource: types.MessageSource{
Sender: types.NewJID("1001", types.DefaultUserServer),
Chat: types.NewJID("1001", types.DefaultUserServer),
},
ID: "mid-img",
PushName: "Bob",
},
Message: &waE2E.Message{
ImageMessage: &waE2E.ImageMessage{
Caption: proto.String("look at this"),
},
},
}
ch.handleIncoming(evt)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
select {
case <-ctx.Done():
t.Fatal("timeout waiting for inbound message")
case inbound, ok := <-messageBus.InboundChan():
if !ok {
t.Fatal("expected inbound message")
}
if inbound.Content != "look at this" {
t.Fatalf("content=%q, want caption only when download fails", inbound.Content)
}
if len(inbound.Media) != 0 {
t.Fatalf("expected no media refs when download fails, got %v", inbound.Media)
}
}
}
func TestHandleIncoming_ImageOnlyNoCaptionDownloadFails_NoInbound(t *testing.T) {
messageBus := bus.NewMessageBus()
ch := &WhatsAppNativeChannel{
BaseChannel: channels.NewBaseChannel("whatsapp_native", config.WhatsAppConfig{}, messageBus, nil),
runCtx: context.Background(),
client: whatsappTestClient(t),
}
evt := &events.Message{
Info: types.MessageInfo{
MessageSource: types.MessageSource{
Sender: types.NewJID("1001", types.DefaultUserServer),
Chat: types.NewJID("1001", types.DefaultUserServer),
},
ID: "mid-img2",
PushName: "Bob",
},
Message: &waE2E.Message{
ImageMessage: &waE2E.ImageMessage{},
},
}
ch.handleIncoming(evt)
ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond)
defer cancel()
select {
case inbound, ok := <-messageBus.InboundChan():
if ok {
t.Fatalf("unexpected inbound when media-only and download fails: %+v", inbound)
}
case <-ctx.Done():
// expected: no message
}
}

View file

@ -11,6 +11,7 @@ import (
"context"
"database/sql"
"fmt"
"mime"
"os"
"path/filepath"
"strings"
@ -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"
)
@ -346,17 +348,119 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
}
senderID := evt.Info.Sender.String()
chatID := evt.Info.Chat.String()
content := evt.Message.GetConversation()
if content == "" && evt.Message.ExtendedTextMessage != nil {
content = evt.Message.ExtendedTextMessage.GetText()
}
content = utils.SanitizeMessageContent(content)
messageID := evt.Info.ID
if content == "" {
sender := bus.SenderInfo{
Platform: "whatsapp",
PlatformID: senderID,
CanonicalID: identity.BuildCanonicalID("whatsapp", senderID),
DisplayName: evt.Info.PushName,
}
if !c.IsAllowedSender(sender) {
return
}
var mediaPaths []string
m := evt.Message
needsClient := m.GetImageMessage() != nil || m.GetVideoMessage() != nil ||
m.GetAudioMessage() != nil || m.GetDocumentMessage() != nil || m.GetStickerMessage() != nil
c.mu.Lock()
client := c.client
c.mu.Unlock()
if needsClient && client == nil {
return
}
content := m.GetConversation()
if content == "" && m.ExtendedTextMessage != nil {
content = m.ExtendedTextMessage.GetText()
}
mediaPaths := []string{}
scope := channels.BuildMediaScope("whatsapp_native", chatID, messageID)
storeMedia := func(localPath, filename string) string {
if store := c.GetMediaStore(); store != nil {
ref, err := store.Store(localPath, media.MediaMeta{
Filename: filename,
Source: "whatsapp_native",
}, scope)
if err == nil {
return ref
}
logger.WarnCF("whatsapp", "Failed to store WhatsApp media in MediaStore", map[string]any{
"path": localPath, "error": err.Error(),
})
}
return localPath
}
appendCaption := func(cap string) {
cap = strings.TrimSpace(cap)
if cap == "" {
return
}
if content != "" {
content += "\n"
}
content += cap
}
appendAnnotation := func(line string) {
if content != "" {
content += "\n"
}
content += line
}
downloadAndStore := func(dm whatsmeow.DownloadableMessage, storeName, annotation string) {
localPath, err := c.downloadWhatsAppMediaToTemp(c.runCtx, client, dm)
if err != nil {
logger.ErrorCF("whatsapp", "Failed to download WhatsApp media", map[string]any{
"error": err.Error(),
})
return
}
mediaPaths = append(mediaPaths, storeMedia(localPath, storeName))
appendAnnotation(annotation)
}
if img := m.GetImageMessage(); img != nil {
appendCaption(img.GetCaption())
fname := imageFilenameFromMime(img.GetMimetype())
downloadAndStore(img, fname, "[image: photo]")
}
if vid := m.GetVideoMessage(); vid != nil {
appendCaption(vid.GetCaption())
downloadAndStore(vid, "video.mp4", "[video]")
}
if aud := m.GetAudioMessage(); aud != nil {
ann := "[audio]"
if aud.GetPTT() {
ann = "[voice]"
}
fname := "audio.ogg"
if mt := aud.GetMimetype(); strings.Contains(mt, "mpeg") || strings.Contains(mt, "mp3") {
fname = "audio.mp3"
}
downloadAndStore(aud, fname, ann)
}
if doc := m.GetDocumentMessage(); doc != nil {
appendCaption(doc.GetCaption())
fname := doc.GetFileName()
if strings.TrimSpace(fname) == "" {
fname = "document"
}
downloadAndStore(doc, fname, "[file: "+fname+"]")
}
if st := m.GetStickerMessage(); st != nil {
downloadAndStore(st, "sticker.webp", "[sticker]")
}
content = utils.SanitizeMessageContent(content)
if content == "" && len(mediaPaths) == 0 {
return
}
metadata := make(map[string]string)
metadata["message_id"] = evt.Info.ID
@ -376,26 +480,51 @@ func (c *WhatsAppNativeChannel) handleIncoming(evt *events.Message) {
peerKind = "group"
}
peer := bus.Peer{Kind: peerKind, ID: chatID}
messageID := evt.Info.ID
sender := bus.SenderInfo{
Platform: "whatsapp",
PlatformID: senderID,
CanonicalID: identity.BuildCanonicalID("whatsapp", senderID),
DisplayName: evt.Info.PushName,
}
if !c.IsAllowedSender(sender) {
return
}
logger.DebugCF(
"whatsapp",
"WhatsApp message received",
map[string]any{"sender_id": senderID, "content_preview": utils.Truncate(content, 50)},
map[string]any{"sender_id": senderID, "content_preview": utils.Truncate(content, 50), "media_count": len(mediaPaths)},
)
c.HandleMessage(c.runCtx, peer, messageID, senderID, chatID, content, mediaPaths, metadata, sender)
}
// downloadWhatsAppMediaToTemp decrypts and writes WhatsApp attachment bytes to a temp file.
func (c *WhatsAppNativeChannel) downloadWhatsAppMediaToTemp(ctx context.Context, cli *whatsmeow.Client, msg whatsmeow.DownloadableMessage) (path string, err error) {
data, err := cli.Download(ctx, msg)
if err != nil {
return "", err
}
f, err := os.CreateTemp("", "whatsapp-native-media-*")
if err != nil {
return "", err
}
path = f.Name()
if _, err = f.Write(data); err != nil {
_ = f.Close()
_ = os.Remove(path)
return "", err
}
if err = f.Close(); err != nil {
_ = os.Remove(path)
return "", err
}
return path, nil
}
func imageFilenameFromMime(mimetype string) string {
switch {
case strings.Contains(mimetype, "png"):
return "photo.png"
case strings.Contains(mimetype, "gif"):
return "photo.gif"
case strings.Contains(mimetype, "webp"):
return "photo.webp"
default:
return "photo.jpg"
}
}
func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessage) error {
if !c.IsRunning() {
return channels.ErrNotRunning
@ -435,6 +564,187 @@ func (c *WhatsAppNativeChannel) Send(ctx context.Context, msg bus.OutboundMessag
return nil
}
// SendMedia implements channels.MediaSender for outbound images, audio, video, and files.
func (c *WhatsAppNativeChannel) SendMedia(ctx context.Context, msg bus.OutboundMediaMessage) error {
if !c.IsRunning() {
return channels.ErrNotRunning
}
select {
case <-ctx.Done():
return ctx.Err()
default:
}
c.mu.Lock()
client := c.client
c.mu.Unlock()
if client == nil || !client.IsConnected() {
return fmt.Errorf("whatsapp connection not established: %w", channels.ErrTemporary)
}
if client.Store.ID == nil {
return fmt.Errorf("whatsapp not yet paired (QR login pending): %w", channels.ErrTemporary)
}
to, err := parseJID(msg.ChatID)
if err != nil {
return fmt.Errorf("invalid chat id %q: %w", msg.ChatID, channels.ErrSendFailed)
}
store := c.GetMediaStore()
if store == nil {
return fmt.Errorf("no media store available: %w", channels.ErrSendFailed)
}
for _, part := range msg.Parts {
localPath, err := store.Resolve(part.Ref)
if err != nil {
logger.ErrorCF("whatsapp", "Failed to resolve media ref for send", map[string]any{
"ref": part.Ref, "error": err.Error(),
})
return fmt.Errorf("resolve media ref: %w", channels.ErrSendFailed)
}
data, err := os.ReadFile(localPath)
if err != nil {
logger.ErrorCF("whatsapp", "Failed to read media file for send", map[string]any{
"path": localPath, "error": err.Error(),
})
return fmt.Errorf("read media: %w", channels.ErrSendFailed)
}
mt := whatsappOutboundMimetype(part, localPath)
waMsg, err := buildWhatsAppMediaMessage(ctx, client, part.Type, data, mt, part.Caption, part.Filename)
if err != nil {
logger.ErrorCF("whatsapp", "Failed to build/upload WhatsApp media message", map[string]any{
"type": part.Type, "error": err.Error(),
})
return fmt.Errorf("whatsapp media upload: %w", channels.ErrTemporary)
}
if _, err = client.SendMessage(ctx, to, waMsg); err != nil {
logger.ErrorCF("whatsapp", "Failed to send WhatsApp media", map[string]any{
"type": part.Type, "error": err.Error(),
})
return fmt.Errorf("whatsapp send media: %w", channels.ErrTemporary)
}
}
return nil
}
func whatsappOutboundMimetype(part bus.MediaPart, localPath string) string {
if part.ContentType != "" {
return part.ContentType
}
ext := filepath.Ext(localPath)
if ext != "" {
if m := mime.TypeByExtension(ext); m != "" {
return m
}
}
switch part.Type {
case "image":
return "image/jpeg"
case "audio":
return "audio/mpeg"
case "video":
return "video/mp4"
default:
return "application/octet-stream"
}
}
func buildWhatsAppMediaMessage(
ctx context.Context,
client *whatsmeow.Client,
partType string,
data []byte,
mimetype, caption, filenameHint string,
) (*waE2E.Message, error) {
switch partType {
case "image":
resp, err := client.Upload(ctx, data, whatsmeow.MediaImage)
if err != nil {
return nil, err
}
img := &waE2E.ImageMessage{
Mimetype: proto.String(mimetype),
URL: proto.String(resp.URL),
DirectPath: proto.String(resp.DirectPath),
MediaKey: resp.MediaKey,
FileEncSHA256: resp.FileEncSHA256,
FileSHA256: resp.FileSHA256,
FileLength: proto.Uint64(resp.FileLength),
}
if caption != "" {
img.Caption = proto.String(caption)
}
return &waE2E.Message{ImageMessage: img}, nil
case "video":
resp, err := client.Upload(ctx, data, whatsmeow.MediaVideo)
if err != nil {
return nil, err
}
vid := &waE2E.VideoMessage{
Mimetype: proto.String(mimetype),
URL: proto.String(resp.URL),
DirectPath: proto.String(resp.DirectPath),
MediaKey: resp.MediaKey,
FileEncSHA256: resp.FileEncSHA256,
FileSHA256: resp.FileSHA256,
FileLength: proto.Uint64(resp.FileLength),
}
if caption != "" {
vid.Caption = proto.String(caption)
}
return &waE2E.Message{VideoMessage: vid}, nil
case "audio":
resp, err := client.Upload(ctx, data, whatsmeow.MediaAudio)
if err != nil {
return nil, err
}
ptt := strings.Contains(mimetype, "ogg") || strings.Contains(mimetype, "opus")
aud := &waE2E.AudioMessage{
Mimetype: proto.String(mimetype),
URL: proto.String(resp.URL),
DirectPath: proto.String(resp.DirectPath),
MediaKey: resp.MediaKey,
FileEncSHA256: resp.FileEncSHA256,
FileSHA256: resp.FileSHA256,
FileLength: proto.Uint64(resp.FileLength),
PTT: proto.Bool(ptt),
}
return &waE2E.Message{AudioMessage: aud}, nil
default: // "file" and unknown
resp, err := client.Upload(ctx, data, whatsmeow.MediaDocument)
if err != nil {
return nil, err
}
fname := strings.TrimSpace(filenameHint)
if fname == "" {
fname = "file"
}
doc := &waE2E.DocumentMessage{
Mimetype: proto.String(mimetype),
FileName: proto.String(fname),
URL: proto.String(resp.URL),
DirectPath: proto.String(resp.DirectPath),
MediaKey: resp.MediaKey,
FileEncSHA256: resp.FileEncSHA256,
FileSHA256: resp.FileSHA256,
FileLength: proto.Uint64(resp.FileLength),
}
if caption != "" {
doc.Caption = proto.String(caption)
}
return &waE2E.Message{DocumentMessage: doc}, nil
}
}
// parseJID converts a chat ID (phone number or JID string) to types.JID.
func parseJID(s string) (types.JID, error) {
s = strings.TrimSpace(s)