diff --git a/cmd/server/main.go b/cmd/server/main.go index 83e6da0..7dc3658 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -37,7 +37,7 @@ func main() { log.Fatalf("mail: %v", err) } bus := events.New() - bot, err := discord.FromEnv(store.NewDiscordLinks(db), bus) + bot, err := discord.FromEnv(store.NewDiscordLinks(db), bus, store.NewPostgres(db), notifier) if err != nil { log.Fatalf("discord: %v", err) } diff --git a/internal/discord/bot.go b/internal/discord/bot.go index cc4bc19..9b04415 100644 --- a/internal/discord/bot.go +++ b/internal/discord/bot.go @@ -13,6 +13,7 @@ import ( "github.com/bwmarrin/discordgo" "plumber/internal/events" + "plumber/internal/mail" "plumber/internal/store" ) @@ -23,6 +24,10 @@ type Bot struct { channelID string links store.DiscordLinkStore api API + store store.Store + mail mail.Notifier + admins map[string]string + botUserID string } // New constructs an outbound subscriber. Tests inject a fake API. @@ -31,7 +36,7 @@ func New(channelID string, links store.DiscordLinkStore, api API) *Bot { } // FromEnv builds a bot when Discord env is set. Missing config is a no-op. -func FromEnv(links store.DiscordLinkStore, bus *events.Bus) (*Bot, error) { +func FromEnv(links store.DiscordLinkStore, bus *events.Bus, st store.Store, mailer mail.Notifier) (*Bot, error) { token := strings.TrimSpace(os.Getenv("DISCORD_BOT_TOKEN")) channelID := strings.TrimSpace(os.Getenv("DISCORD_CHANNEL_ID")) if token == "" && channelID == "" { @@ -50,11 +55,26 @@ func FromEnv(links store.DiscordLinkStore, bus *events.Bus) (*Bot, error) { if err != nil { return nil, err } + session.Identify.Intents = discordgo.IntentsGuilds | discordgo.IntentsGuildMessages | discordgo.IntentsMessageContent + if mailer == nil { + mailer = mail.Nop{} + } bot := New(channelID, links, &sessionAPI{session: session}) + bot.store = st + bot.mail = mailer + bot.admins = parseAdminMap(os.Getenv("DISCORD_ADMIN_MAP")) + session.AddHandler(bot.onMessageCreate) if bus != nil { bus.Subscribe(bot.Handle) } - log.Printf("discord: outbound subscriber enabled") + if err := session.Open(); err != nil { + _ = session.Close() + return nil, fmt.Errorf("discord gateway: %w", err) + } + if session.State != nil && session.State.User != nil { + bot.botUserID = session.State.User.ID + } + log.Printf("discord: subscriber enabled") return bot, nil } diff --git a/internal/discord/bot_test.go b/internal/discord/bot_test.go index 2c5e178..f6381db 100644 --- a/internal/discord/bot_test.go +++ b/internal/discord/bot_test.go @@ -204,7 +204,7 @@ func TestFormatMessage(t *testing.T) { func TestFromEnvDisabled(t *testing.T) { t.Setenv("DISCORD_BOT_TOKEN", "") t.Setenv("DISCORD_CHANNEL_ID", "") - bot, err := FromEnv(newMemoryLinks(), nil) + bot, err := FromEnv(newMemoryLinks(), nil, nil, nil) if err != nil || bot != nil { t.Fatalf("disabled FromEnv = (%v, %v)", bot, err) } @@ -213,12 +213,12 @@ func TestFromEnvDisabled(t *testing.T) { func TestFromEnvRequiresBoth(t *testing.T) { t.Setenv("DISCORD_BOT_TOKEN", "token") t.Setenv("DISCORD_CHANNEL_ID", "") - if _, err := FromEnv(newMemoryLinks(), nil); err == nil { + if _, err := FromEnv(newMemoryLinks(), nil, nil, nil); err == nil { t.Fatal("expected error when channel is missing") } t.Setenv("DISCORD_BOT_TOKEN", "") t.Setenv("DISCORD_CHANNEL_ID", "channel") - if _, err := FromEnv(newMemoryLinks(), nil); err == nil { + if _, err := FromEnv(newMemoryLinks(), nil, nil, nil); err == nil { t.Fatal("expected error when token is missing") } } diff --git a/internal/discord/inbound.go b/internal/discord/inbound.go new file mode 100644 index 0000000..c7414a7 --- /dev/null +++ b/internal/discord/inbound.go @@ -0,0 +1,227 @@ +package discord + +import ( + "context" + "database/sql" + "errors" + "fmt" + "log" + "strings" + "time" + + "github.com/bwmarrin/discordgo" + + "plumber/internal/mail" + "plumber/internal/store" +) + +const inboundBodyLimit = 12000 + +type inboundMessage struct { + ID string + ChannelID string + GuildID string + AuthorID string + Content string + ReferencedMessageID string + Bot bool + Attachments int +} + +func parseAdminMap(raw string) map[string]string { + out := map[string]string{} + for _, part := range strings.Split(raw, ",") { + part = strings.TrimSpace(part) + if part == "" { + continue + } + id, username, ok := strings.Cut(part, ":") + id = strings.TrimSpace(id) + username = store.NormalizeUsername(username) + if !ok || id == "" || username == "" { + log.Printf("discord: skip invalid DISCORD_ADMIN_MAP entry %q", part) + continue + } + out[id] = username + } + return out +} + +func (b *Bot) onMessageCreate(_ *discordgo.Session, m *discordgo.MessageCreate) { + if b == nil || m == nil || m.Author == nil { + return + } + in := inboundMessage{ + ID: m.ID, + ChannelID: m.ChannelID, + GuildID: m.GuildID, + AuthorID: m.Author.ID, + Content: m.Content, + Bot: m.Author.Bot, + Attachments: len(m.Attachments), + } + if m.MessageReference != nil { + in.ReferencedMessageID = m.MessageReference.MessageID + } + b.handleInbound(in) +} + +func (b *Bot) handleInbound(in inboundMessage) { + if b == nil || b.store == nil { + return + } + if in.Bot || strings.TrimSpace(in.GuildID) == "" { + return + } + if b.botUserID != "" && in.AuthorID == b.botUserID { + return + } + body := strings.TrimSpace(in.Content) + if in.Attachments > 0 { + log.Printf("discord: ignoring %d attachment(s) on %s", in.Attachments, in.ID) + } + if body == "" { + return + } + ctx, cancel := context.WithTimeout(context.Background(), discordTimeout) + defer cancel() + if !b.knownChannel(ctx, in.ChannelID) { + return + } + username := b.admins[in.AuthorID] + if username == "" { + return + } + author, err := b.store.UserByUsername(ctx, username) + if err != nil { + if !errors.Is(err, sql.ErrNoRows) { + log.Printf("discord: inbound author %s: %v", username, err) + } + return + } + if !author.Admin() { + log.Printf("discord: inbound %s is not an admin", username) + return + } + parent, root, err := b.inboundParent(ctx, in) + if err != nil { + if !errors.Is(err, sql.ErrNoRows) { + log.Printf("discord: inbound parent %s: %v", in.ID, err) + } + return + } + if root.PostState == store.PostStateHidden { + return + } + reply := &store.Post{ + AuthorID: author.ID, + Body: truncateRunes(body, inboundBodyLimit), + ParentID: &parent.ID, + } + if err := b.store.CreatePost(ctx, reply); err != nil { + log.Printf("discord: create inbound %s: %v", in.ID, err) + return + } + if err := b.links.Upsert(ctx, store.DiscordLink{ + PostID: reply.ID, + MessageID: in.ID, + }); err != nil { + log.Printf("discord: save inbound link %s: %v", reply.ID, err) + } + b.notifyInboundReply(parent, root, reply, author) + log.Printf("discord: inbound reply %s -> post %s", in.ID, reply.ID) +} + +func (b *Bot) knownChannel(ctx context.Context, channelID string) bool { + if strings.TrimSpace(channelID) == "" { + return false + } + if channelID == b.channelID { + return true + } + _, err := b.links.GetRootByThreadID(ctx, channelID) + return err == nil +} + +func (b *Bot) inboundParent(ctx context.Context, in inboundMessage) (*store.Post, *store.Post, error) { + if ref := strings.TrimSpace(in.ReferencedMessageID); ref != "" { + link, err := b.links.GetByMessageID(ctx, ref) + if err == nil { + return b.postAndRoot(ctx, link.PostID) + } + if !errors.Is(err, sql.ErrNoRows) { + return nil, nil, err + } + } + link, err := b.links.GetRootByThreadID(ctx, in.ChannelID) + if err != nil { + return nil, nil, err + } + return b.postAndRoot(ctx, link.PostID) +} + +func (b *Bot) postAndRoot(ctx context.Context, postID string) (*store.Post, *store.Post, error) { + postID = strings.TrimSpace(postID) + if postID == "" { + return nil, nil, sql.ErrNoRows + } + post, err := b.store.GetPost(ctx, postID) + if err != nil { + return nil, nil, err + } + current := post + seen := map[string]bool{} + for current.ParentID != nil { + if seen[current.ID] { + return nil, nil, fmt.Errorf("post ancestry cycle at %s", current.ID) + } + seen[current.ID] = true + current, err = b.store.GetPost(ctx, *current.ParentID) + if err != nil { + return nil, nil, err + } + } + return post, current, nil +} + +func (b *Bot) notifyInboundReply(parent, root, reply *store.Post, author *store.User) { + if parent == nil || root == nil || reply == nil || author == nil || b.mail == nil { + return + } + if _, disabled := b.mail.(mail.Nop); disabled { + return + } + recipientID := parent.AuthorID + if author.Admin() { + recipientID = root.AuthorID + } + if recipientID == author.ID { + return + } + msg := mail.PostReply{ + RootID: root.ID, + RootTitle: root.Title, + ReplyID: reply.ID, + ReplyBody: reply.Body, + ReplyAuthorName: author.Name, + } + go func() { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + recipient, err := b.store.UserByID(ctx, recipientID) + if err != nil { + log.Printf("notify reply %s: load recipient: %v", msg.ReplyID, err) + return + } + if recipient == nil || strings.TrimSpace(recipient.Email) == "" { + return + } + msg.ToEmail = recipient.Email + msg.ToName = recipient.Name + if err := b.mail.NotifyPostReply(ctx, msg); err != nil { + log.Printf("notify reply %s: %v", msg.ReplyID, err) + return + } + log.Printf("notify reply %s: accepted", msg.ReplyID) + }() +} diff --git a/internal/discord/inbound_test.go b/internal/discord/inbound_test.go new file mode 100644 index 0000000..53d925b --- /dev/null +++ b/internal/discord/inbound_test.go @@ -0,0 +1,218 @@ +package discord + +import ( + "context" + "testing" + "time" + + "plumber/internal/events" + "plumber/internal/mail" + "plumber/internal/pacific" + "plumber/internal/store" +) + +func TestParseAdminMap(t *testing.T) { + t.Parallel() + + got := parseAdminMap(" 123:Plumber ,456:other,bad, :empty,789: ") + if got["123"] != "plumber" || got["456"] != "other" || len(got) != 2 { + t.Fatalf("parseAdminMap = %#v", got) + } +} + +func TestInboundCreatesSiteReply(t *testing.T) { + t.Parallel() + + mem, homeowner, admin := seedInboundUsers(t) + links := newMemoryLinks() + api := &fakeAPI{} + mailer := &mail.Recording{} + bot := inboundTestBot(mem, links, api, mailer, admin.Username) + bus := events.New() + defer bus.Close() + bus.Subscribe(bot.Handle) + root := seedLinkedRoot(t, mem, links, homeowner.ID, "thread-1") + + bot.handleInbound(inboundMessage{ + ID: "d-reply-1", + ChannelID: "thread-1", + GuildID: "guild-1", + AuthorID: "snow-admin", + Content: "Replace the cartridge.", + }) + + thread, err := mem.GetPostThread(context.Background(), root.ID) + if err != nil || len(thread.Replies) != 1 { + t.Fatalf("thread = %+v, %v", thread, err) + } + reply := thread.Replies[0] + if reply.AuthorID != admin.ID || reply.Body != "Replace the cartridge." || reply.ParentID == nil || *reply.ParentID != root.ID { + t.Fatalf("reply = %+v", reply) + } + link, err := links.GetByPostID(context.Background(), reply.ID) + if err != nil || link.MessageID != "d-reply-1" || link.ThreadID != "" { + t.Fatalf("inbound link = %+v, %v", link, err) + } + time.Sleep(20 * time.Millisecond) + if len(api.sends) != 0 || len(api.edits) != 0 { + t.Fatalf("inbound echoed to Discord: sends=%+v edits=%+v", api.sends, api.edits) + } + msgs := waitForMail(t, mailer, 1) + if msgs[0].ToEmail != homeowner.Email || msgs[0].ReplyID != reply.ID || msgs[0].RootID != root.ID { + t.Fatalf("mail = %+v", msgs[0]) + } +} + +func TestInboundParentsFromReference(t *testing.T) { + t.Parallel() + + mem, homeowner, admin := seedInboundUsers(t) + links := newMemoryLinks() + bot := inboundTestBot(mem, links, &fakeAPI{}, &mail.Recording{}, admin.Username) + root := seedLinkedRoot(t, mem, links, homeowner.ID, "thread-1") + plumberReply := &store.Post{ParentID: &root.ID, AuthorID: admin.ID, Body: "First look."} + if err := mem.CreatePost(context.Background(), plumberReply); err != nil { + t.Fatal(err) + } + if err := links.Upsert(context.Background(), store.DiscordLink{ + PostID: plumberReply.ID, + MessageID: "d-plumber-1", + }); err != nil { + t.Fatal(err) + } + + bot.handleInbound(inboundMessage{ + ID: "d-nested", + ChannelID: "thread-1", + GuildID: "guild-1", + AuthorID: "snow-admin", + Content: "More detail.", + ReferencedMessageID: "d-plumber-1", + }) + + thread, err := mem.GetPostThread(context.Background(), root.ID) + if err != nil || len(thread.Replies) != 1 || len(thread.Replies[0].Replies) != 1 { + t.Fatalf("thread = %+v, %v", thread, err) + } + nested := thread.Replies[0].Replies[0] + if nested.ParentID == nil || *nested.ParentID != plumberReply.ID { + t.Fatalf("nested parent = %+v", nested) + } +} + +func TestInboundIgnoresAllowlistHiddenAndEchoSources(t *testing.T) { + t.Parallel() + + mem, homeowner, admin := seedInboundUsers(t) + links := newMemoryLinks() + api := &fakeAPI{} + bot := inboundTestBot(mem, links, api, mail.Nop{}, admin.Username) + root := seedLinkedRoot(t, mem, links, homeowner.ID, "thread-1") + hidden := &store.Post{ + AuthorID: homeowner.ID, + Title: "Hidden", + Body: "No.", + PostDate: pacific.Today(), + PostState: store.PostStateHidden, + } + if err := mem.CreatePost(context.Background(), hidden); err != nil { + t.Fatal(err) + } + if err := links.Upsert(context.Background(), store.DiscordLink{ + PostID: hidden.ID, + MessageID: "d-hidden", + ThreadID: "thread-hidden", + }); err != nil { + t.Fatal(err) + } + + cases := []inboundMessage{ + {ID: "bot", ChannelID: "thread-1", GuildID: "g", AuthorID: "snow-admin", Content: "x", Bot: true}, + {ID: "dm", ChannelID: "thread-1", AuthorID: "snow-admin", Content: "x"}, + {ID: "self", ChannelID: "thread-1", GuildID: "g", AuthorID: "bot-1", Content: "x"}, + {ID: "stranger", ChannelID: "thread-1", GuildID: "g", AuthorID: "snow-other", Content: "x"}, + {ID: "elsewhere", ChannelID: "other-thread", GuildID: "g", AuthorID: "snow-admin", Content: "x"}, + {ID: "empty", ChannelID: "thread-1", GuildID: "g", AuthorID: "snow-admin", Content: " ", Attachments: 1}, + {ID: "hidden", ChannelID: "thread-hidden", GuildID: "g", AuthorID: "snow-admin", Content: "x"}, + {ID: "channel-root", ChannelID: "channel-1", GuildID: "g", AuthorID: "snow-admin", Content: "new question"}, + } + for _, in := range cases { + bot.handleInbound(in) + } + thread, err := mem.GetPostThread(context.Background(), root.ID) + if err != nil || len(thread.Replies) != 0 { + t.Fatalf("unexpected replies: %+v, %v", thread, err) + } + if len(api.sends) != 0 { + t.Fatalf("unexpected discord sends %+v", api.sends) + } +} + +func inboundTestBot(mem *store.Memory, links *memoryLinks, api *fakeAPI, mailer mail.Notifier, adminUsername string) *Bot { + bot := New("channel-1", links, api) + bot.store = mem + bot.mail = mailer + bot.admins = map[string]string{"snow-admin": adminUsername} + bot.botUserID = "bot-1" + return bot +} + +func seedInboundUsers(t *testing.T) (*store.Memory, *store.User, *store.User) { + t.Helper() + mem := store.NewMemory() + homeowner := &store.User{ + Username: "homeowner", + Name: "Sam", + Email: "sam@example.com", + PasswordHash: "x", + Role: store.RoleUser, + } + if err := mem.CreateUser(context.Background(), homeowner); err != nil { + t.Fatal(err) + } + admin := &store.User{ + Username: "plumber", + Name: "Pat", + Email: "pat@example.com", + PasswordHash: "x", + Role: store.RoleAdmin, + } + if err := mem.CreateUser(context.Background(), admin); err != nil { + t.Fatal(err) + } + return mem, homeowner, admin +} + +func seedLinkedRoot(t *testing.T, mem *store.Memory, links *memoryLinks, authorID, threadID string) *store.Post { + t.Helper() + root := &store.Post{ + AuthorID: authorID, + Title: "Leaky sink", + Body: "It drips.", + PostDate: pacific.Today(), + } + if err := mem.CreatePost(context.Background(), root); err != nil { + t.Fatal(err) + } + if err := links.Upsert(context.Background(), store.DiscordLink{ + PostID: root.ID, + MessageID: "d-root", + ThreadID: threadID, + }); err != nil { + t.Fatal(err) + } + return root +} + +func waitForMail(t *testing.T, recording *mail.Recording, want int) []mail.PostReply { + t.Helper() + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + if recording.Len() >= want { + return recording.Snapshot() + } + time.Sleep(10 * time.Millisecond) + } + t.Fatalf("recorded %d notifications, want %d", recording.Len(), want) + return nil +}