diff --git a/internal/discord/api.go b/internal/discord/api.go index 58d0bda..ab964ef 100644 --- a/internal/discord/api.go +++ b/internal/discord/api.go @@ -9,7 +9,7 @@ import ( // API is the Discord REST surface used by the outbound subscriber. type API interface { SendToChannel(ctx context.Context, channelID string, msg Message) (messageID string, err error) - StartThread(ctx context.Context, channelID, messageID, name string) (threadID string, err error) + StartThread(ctx context.Context, channelID, name string) (threadID string, err error) SendToThread(ctx context.Context, threadID string, msg Message) (messageID string, err error) Edit(ctx context.Context, channelID, messageID string, msg Message) error Close() error @@ -27,9 +27,10 @@ func (s *sessionAPI) SendToChannel(_ context.Context, channelID string, msg Mess return sent.ID, nil } -func (s *sessionAPI) StartThread(_ context.Context, channelID, messageID, name string) (string, error) { - thread, err := s.session.MessageThreadStartComplex(channelID, messageID, &discordgo.ThreadStart{ +func (s *sessionAPI) StartThread(_ context.Context, channelID, name string) (string, error) { + thread, err := s.session.ThreadStartComplex(channelID, &discordgo.ThreadStart{ Name: name, + Type: discordgo.ChannelTypeGuildPublicThread, AutoArchiveDuration: 10080, }) if err != nil { @@ -43,10 +44,12 @@ func (s *sessionAPI) SendToThread(ctx context.Context, threadID string, msg Mess } func (s *sessionAPI) Edit(_ context.Context, channelID, messageID string, msg Message) error { + content := messageContent(msg) embeds := toEmbeds(msg) _, err := s.session.ChannelMessageEditComplex(&discordgo.MessageEdit{ ID: messageID, Channel: channelID, + Content: &content, Embeds: &embeds, }) return err @@ -61,6 +64,7 @@ func (s *sessionAPI) Close() error { func toMessageSend(msg Message) *discordgo.MessageSend { return &discordgo.MessageSend{ + Content: messageContent(msg), Embeds: toEmbeds(msg), AllowedMentions: &discordgo.MessageAllowedMentions{}, } @@ -69,7 +73,7 @@ func toMessageSend(msg Message) *discordgo.MessageSend { func toEmbeds(msg Message) []*discordgo.MessageEmbed { main := &discordgo.MessageEmbed{ Title: msg.Title, - URL: msg.URL, + URL: publicURL(msg.URL), Description: msg.Description, Color: embedColor, } diff --git a/internal/discord/bot.go b/internal/discord/bot.go index 9b04415..b127969 100644 --- a/internal/discord/bot.go +++ b/internal/discord/bot.go @@ -133,16 +133,16 @@ func (b *Bot) onUpdated(ctx context.Context, ev events.PostEvent) { func (b *Bot) createRoot(ctx context.Context, ev events.PostEvent) { msg := formatMessage(ev) - messageID, err := b.api.SendToChannel(ctx, b.channelID, msg) - if err != nil { - log.Printf("discord: send root %s: %v", ev.PostID, err) - return - } - threadID, err := b.api.StartThread(ctx, b.channelID, messageID, msg.ThreadName) + threadID, err := b.api.StartThread(ctx, b.channelID, msg.ThreadName) if err != nil { log.Printf("discord: start thread %s: %v", ev.PostID, err) return } + messageID, err := b.api.SendToThread(ctx, threadID, msg) + if err != nil { + log.Printf("discord: send root %s: %v", ev.PostID, err) + return + } if err := b.links.Upsert(ctx, store.DiscordLink{ PostID: ev.PostID, MessageID: messageID, @@ -185,7 +185,7 @@ func (b *Bot) createReply(ctx context.Context, ev events.PostEvent) { func (b *Bot) editChannel(ctx context.Context, ev events.PostEvent, link *store.DiscordLink) (string, error) { if strings.TrimSpace(link.ThreadID) != "" { - return b.channelID, nil + return link.ThreadID, nil } root, err := b.links.GetByPostID(ctx, ev.RootID) if err != nil { diff --git a/internal/discord/bot_test.go b/internal/discord/bot_test.go index f6381db..fde963b 100644 --- a/internal/discord/bot_test.go +++ b/internal/discord/bot_test.go @@ -2,11 +2,11 @@ package discord import ( "context" + "strconv" + "strings" "sync" "testing" - "strconv" - "plumber/internal/events" "plumber/internal/store" ) @@ -30,17 +30,18 @@ func (f *fakeAPI) SendToChannel(_ context.Context, channelID string, msg Message return f.record("channel", channelID, "", msg) } -func (f *fakeAPI) StartThread(_ context.Context, channelID, messageID, name string) (string, error) { +func (f *fakeAPI) StartThread(_ context.Context, channelID, name string) (string, error) { f.mu.Lock() defer f.mu.Unlock() f.next++ + id := "thread-" + strconv.Itoa(f.next) f.sends = append(f.sends, recordedSend{ Kind: "thread", ChannelID: channelID, Name: name, - Msg: Message{ThreadName: name, URL: messageID}, + Msg: Message{ThreadName: name}, }) - return "thread-" + messageID, nil + return id, nil } func (f *fakeAPI) SendToThread(_ context.Context, threadID string, msg Message) (string, error) { @@ -93,17 +94,21 @@ func TestOutboundRootReplyAndEdit(t *testing.T) { } bot.Handle(ctx, events.PostCreated{PostEvent: root}) - if len(api.sends) != 2 || api.sends[0].Kind != "channel" || api.sends[1].Kind != "thread" { + if len(api.sends) != 2 || + api.sends[0].Kind != "thread" || + api.sends[1].Kind != "thread-msg" { t.Fatalf("root sends = %+v", api.sends) } - if api.sends[0].ChannelID != "channel-1" || api.sends[1].Name != "Leaky sink" { + if api.sends[0].ChannelID != "channel-1" || + api.sends[0].Name != "sam asks: Leaky sink" || + api.sends[1].ChannelID != "thread-1" { t.Fatalf("root routing = %+v", api.sends) } - if got := api.sends[0].Msg.ImageURLs; len(got) != 2 || got[0] != "https://cdn.example/a.jpg" { + if got := api.sends[1].Msg.ImageURLs; len(got) != 2 || got[0] != "https://cdn.example/a.jpg" { t.Fatalf("root images = %v", got) } link, err := links.GetByPostID(ctx, "root-1") - if err != nil || link.MessageID != "msg-1" || link.ThreadID != "thread-msg-1" { + if err != nil || link.MessageID != "msg-2" || link.ThreadID != "thread-1" { t.Fatalf("root link = %+v, %v", link, err) } @@ -116,7 +121,7 @@ func TestOutboundRootReplyAndEdit(t *testing.T) { Permalink: "https://www.askaplumberfirst.com/questions/root-1#post-reply-1", } bot.Handle(ctx, events.PostCreated{PostEvent: reply}) - if len(api.sends) != 3 || api.sends[2].Kind != "thread-msg" || api.sends[2].ChannelID != "thread-msg-1" { + if len(api.sends) != 3 || api.sends[2].Kind != "thread-msg" || api.sends[2].ChannelID != "thread-1" { t.Fatalf("reply sends = %+v", api.sends) } replyLink, err := links.GetByPostID(ctx, "reply-1") @@ -126,7 +131,7 @@ func TestOutboundRootReplyAndEdit(t *testing.T) { root.Body = "Updated leak." bot.Handle(ctx, events.PostUpdated{PostEvent: root}) - if len(api.edits) != 1 || api.edits[0].ChannelID != "channel-1" || api.edits[0].Name != "msg-1" { + if len(api.edits) != 1 || api.edits[0].ChannelID != "thread-1" || api.edits[0].Name != "msg-2" { t.Fatalf("root edit = %+v", api.edits) } if api.edits[0].Msg.Description != "Updated leak." { @@ -135,7 +140,7 @@ func TestOutboundRootReplyAndEdit(t *testing.T) { reply.Body = "Use a ceramic cartridge." bot.Handle(ctx, events.PostUpdated{PostEvent: reply}) - if len(api.edits) != 2 || api.edits[1].ChannelID != "thread-msg-1" || api.edits[1].Name != "msg-3" { + if len(api.edits) != 2 || api.edits[1].ChannelID != "thread-1" || api.edits[1].Name != "msg-3" { t.Fatalf("reply edit = %+v", api.edits) } } @@ -190,15 +195,24 @@ func TestFormatMessage(t *testing.T) { got.City != "Oakland" || got.Author != "sam" || got.URL != "https://example.com/q" || - got.ThreadName != "Leaky sink" || + got.ThreadName != "sam asks: Leaky sink" || len(got.ImageURLs) != 1 { t.Fatalf("format = %+v", got) } reply := formatMessage(events.PostEvent{Body: "Thanks", AuthorName: ""}) - if reply.Title != "Reply" || reply.Author != "Someone" || reply.ThreadName != "Question" { + if reply.Title != "Reply" || reply.Author != "Someone" || reply.ThreadName != "Someone asks: Question" { t.Fatalf("reply format = %+v", reply) } + content := messageContent(got) + if strings.Contains(content, "Leaky sink") || + !strings.Contains(content, "It drips.") || + !strings.Contains(content, "Oakland") { + t.Fatalf("content = %q", content) + } + if publicURL("/questions/x") != "" || publicURL("http://localhost:8080/q") != "" { + t.Fatal("localhost or relative permalink should not be an embed URL") + } } func TestFromEnvDisabled(t *testing.T) { diff --git a/internal/discord/format.go b/internal/discord/format.go index f552a19..c14617a 100644 --- a/internal/discord/format.go +++ b/internal/discord/format.go @@ -39,7 +39,7 @@ func formatMessage(ev events.PostEvent) Message { Description: truncateRunes(strings.TrimSpace(ev.Body), embedDescriptionLimit), City: strings.TrimSpace(ev.City), Author: author, - ThreadName: threadName(ev.Title), + ThreadName: threadName(author, ev.Title), } for _, img := range ev.Images { url := strings.TrimSpace(img.URL) @@ -51,12 +51,16 @@ func formatMessage(ev events.PostEvent) Message { return msg } -func threadName(title string) string { +func threadName(author, title string) string { + author = strings.TrimSpace(author) + if author == "" { + author = "Someone" + } title = strings.TrimSpace(title) if title == "" { - return "Question" + title = "Question" } - return truncateRunes(title, threadNameLimit) + return truncateRunes(author+" asks: "+title, threadNameLimit) } func truncateRunes(s string, max int) string { @@ -73,3 +77,35 @@ func truncateRunes(s string, max int) string { func isRoot(ev events.PostEvent) bool { return strings.TrimSpace(ev.ParentID) == "" } + +func messageContent(msg Message) string { + var parts []string + if body := strings.TrimSpace(msg.Description); body != "" { + parts = append(parts, body) + } + var meta []string + if msg.City != "" { + meta = append(meta, msg.City) + } + if msg.Author != "" { + meta = append(meta, msg.Author) + } + if len(meta) > 0 { + parts = append(parts, strings.Join(meta, " ยท ")) + } + if u := publicURL(msg.URL); u != "" { + parts = append(parts, u) + } + return truncateRunes(strings.Join(parts, "\n"), 2000) +} + +func publicURL(raw string) string { + raw = strings.TrimSpace(raw) + if !strings.HasPrefix(raw, "https://") { + return "" + } + if strings.Contains(raw, "localhost") || strings.Contains(raw, "127.0.0.1") { + return "" + } + return raw +}