diff --git a/.env.example b/.env.example index 0a4207b..2f35705 100644 --- a/.env.example +++ b/.env.example @@ -26,3 +26,7 @@ SECURE_COOKIE=0 # SPACES_BUCKET=your-bucket # SPACES_ENDPOINT=https://nyc3.digitaloceanspaces.com # SPACES_CDN_BASE=https://your-bucket.nyc3.cdn.digitaloceanspaces.com +# Discord bot subscriber. Leave unset to disable. +# DISCORD_BOT_TOKEN= +# DISCORD_CHANNEL_ID= +# DISCORD_ADMIN_MAP=123456789012345678:plumber,234567890123456789:otheradmin diff --git a/cmd/server/main.go b/cmd/server/main.go index 37f986e..83e6da0 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -17,6 +17,7 @@ import ( "plumber" "plumber/internal/blob" + "plumber/internal/discord" "plumber/internal/events" "plumber/internal/mail" "plumber/internal/store" @@ -35,7 +36,15 @@ func main() { if err != nil { log.Fatalf("mail: %v", err) } - handler := newHandler(db, sessions, uploader, notifier, events.New()) + bus := events.New() + bot, err := discord.FromEnv(store.NewDiscordLinks(db), bus) + if err != nil { + log.Fatalf("discord: %v", err) + } + if bot != nil { + defer bot.Close() + } + handler := newHandler(db, sessions, uploader, notifier, bus) run(&http.Server{ Addr: listenAddr(), Handler: handler, diff --git a/db/queries/discord_links.sql b/db/queries/discord_links.sql new file mode 100644 index 0000000..9ac2607 --- /dev/null +++ b/db/queries/discord_links.sql @@ -0,0 +1,32 @@ +-- name: GetDiscordPostLinkByPostID :one +SELECT post_id, discord_message_id, discord_thread_id, created_at +FROM discord_post_links +WHERE post_id = sqlc.arg(post_id); + +-- name: GetDiscordPostLinkByMessageID :one +SELECT post_id, discord_message_id, discord_thread_id, created_at +FROM discord_post_links +WHERE discord_message_id = sqlc.arg(discord_message_id); + +-- name: GetDiscordPostLinkByThreadID :one +SELECT post_id, discord_message_id, discord_thread_id, created_at +FROM discord_post_links +WHERE discord_thread_id = sqlc.arg(discord_thread_id) + AND discord_thread_id <> ''; + +-- name: UpsertDiscordPostLink :exec +INSERT INTO discord_post_links ( + post_id, discord_message_id, discord_thread_id, created_at +) +VALUES ( + sqlc.arg(post_id), + sqlc.arg(discord_message_id), + sqlc.arg(discord_thread_id), + sqlc.arg(created_at) +) +ON CONFLICT (post_id) DO UPDATE SET + discord_message_id = EXCLUDED.discord_message_id, + discord_thread_id = CASE + WHEN EXCLUDED.discord_thread_id <> '' THEN EXCLUDED.discord_thread_id + ELSE discord_post_links.discord_thread_id + END; diff --git a/go.mod b/go.mod index 2ff2b4b..7bb789c 100644 --- a/go.mod +++ b/go.mod @@ -7,10 +7,12 @@ require ( github.com/aws/aws-sdk-go-v2 v1.43.7 github.com/aws/aws-sdk-go-v2/credentials v1.19.37 github.com/aws/aws-sdk-go-v2/service/s3 v1.107.3 + github.com/bwmarrin/discordgo v0.29.0 github.com/go-chi/chi/v5 v5.3.1 github.com/google/uuid v1.6.0 github.com/jackc/pgx/v5 v5.10.0 github.com/joho/godotenv v1.5.1 + github.com/resend/resend-go/v3 v3.16.0 golang.org/x/crypto v0.55.0 golang.org/x/image v0.45.0 ) @@ -25,10 +27,11 @@ require ( github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.38 // indirect github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.39 // indirect github.com/aws/smithy-go v1.27.8 // indirect + github.com/gorilla/websocket v1.4.2 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect github.com/jackc/puddle/v2 v2.2.2 // indirect - github.com/resend/resend-go/v3 v3.16.0 // indirect golang.org/x/sync v0.22.0 // indirect + golang.org/x/sys v0.47.0 // indirect golang.org/x/text v0.41.0 // indirect ) diff --git a/go.sum b/go.sum index 479e538..1ce8c5f 100644 --- a/go.sum +++ b/go.sum @@ -24,6 +24,8 @@ github.com/aws/aws-sdk-go-v2/service/s3 v1.107.3 h1:IKoCZqfWfZzSBi16QFQ+QcbQ3LRQ github.com/aws/aws-sdk-go-v2/service/s3 v1.107.3/go.mod h1:RBpRcXiM4s2pOInVs32GsBonnje+fiAj4mcrStRmlCA= github.com/aws/smithy-go v1.27.8 h1:FR0dxZfIlV7Z8eh2iHfIofdunw382XsDV3Mxt9nUvRY= github.com/aws/smithy-go v1.27.8/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc= +github.com/bwmarrin/discordgo v0.29.0 h1:FmWeXFaKUwrcL3Cx65c20bTRW+vOb6k8AnaP+EgjDno= +github.com/bwmarrin/discordgo v0.29.0/go.mod h1:NJZpH+1AfhIcyQsPeuBKsUtYrRnjkyu0kIVMCHkZtRY= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= @@ -31,6 +33,8 @@ github.com/go-chi/chi/v5 v5.3.1 h1:3j4HZLGZQ3JpMCrPJF/Jl3mYJfWLKBfNJ6quurUGCf8= github.com/go-chi/chi/v5 v5.3.1/go.mod h1:R+tYY2hNuVUUjxoPtqUdgBqevM9s9njzkTLutVsOCto= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/gorilla/websocket v1.4.2 h1:+/TMaTYc4QFitKJxsQ7Yye35DkWvkdLcvGKqM+x0Ufc= +github.com/gorilla/websocket v1.4.2/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= @@ -50,14 +54,22 @@ github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UV github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +golang.org/x/crypto v0.0.0-20210421170649-83a5a9bb288b/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4= golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= golang.org/x/image v0.45.0 h1:FMb1nTbH5H9vF55SriQHgFw5GnNL9Jg6L25BwXKzhB0= golang.org/x/image v0.45.0/go.mod h1:n62x/7RqlwXDvGsSU4u6IUTUf6KghUZ9Bt7cG/T9Fx4= +golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= diff --git a/internal/discord/api.go b/internal/discord/api.go new file mode 100644 index 0000000..58d0bda --- /dev/null +++ b/internal/discord/api.go @@ -0,0 +1,102 @@ +package discord + +import ( + "context" + + "github.com/bwmarrin/discordgo" +) + +// 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) + SendToThread(ctx context.Context, threadID string, msg Message) (messageID string, err error) + Edit(ctx context.Context, channelID, messageID string, msg Message) error + Close() error +} + +type sessionAPI struct { + session *discordgo.Session +} + +func (s *sessionAPI) SendToChannel(_ context.Context, channelID string, msg Message) (string, error) { + sent, err := s.session.ChannelMessageSendComplex(channelID, toMessageSend(msg)) + if err != nil { + return "", err + } + 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{ + Name: name, + AutoArchiveDuration: 10080, + }) + if err != nil { + return "", err + } + return thread.ID, nil +} + +func (s *sessionAPI) SendToThread(ctx context.Context, threadID string, msg Message) (string, error) { + return s.SendToChannel(ctx, threadID, msg) +} + +func (s *sessionAPI) Edit(_ context.Context, channelID, messageID string, msg Message) error { + embeds := toEmbeds(msg) + _, err := s.session.ChannelMessageEditComplex(&discordgo.MessageEdit{ + ID: messageID, + Channel: channelID, + Embeds: &embeds, + }) + return err +} + +func (s *sessionAPI) Close() error { + if s == nil || s.session == nil { + return nil + } + return s.session.Close() +} + +func toMessageSend(msg Message) *discordgo.MessageSend { + return &discordgo.MessageSend{ + Embeds: toEmbeds(msg), + AllowedMentions: &discordgo.MessageAllowedMentions{}, + } +} + +func toEmbeds(msg Message) []*discordgo.MessageEmbed { + main := &discordgo.MessageEmbed{ + Title: msg.Title, + URL: msg.URL, + Description: msg.Description, + Color: embedColor, + } + if msg.City != "" { + main.Fields = append(main.Fields, &discordgo.MessageEmbedField{ + Name: "City", + Value: msg.City, + Inline: true, + }) + } + if msg.Author != "" { + main.Fields = append(main.Fields, &discordgo.MessageEmbedField{ + Name: "Author", + Value: msg.Author, + Inline: true, + }) + } + embeds := []*discordgo.MessageEmbed{main} + for i, url := range msg.ImageURLs { + if i == 0 { + main.Image = &discordgo.MessageEmbedImage{URL: url} + continue + } + embeds = append(embeds, &discordgo.MessageEmbed{ + Color: embedColor, + Image: &discordgo.MessageEmbedImage{URL: url}, + }) + } + return embeds +} diff --git a/internal/discord/bot.go b/internal/discord/bot.go new file mode 100644 index 0000000..cc4bc19 --- /dev/null +++ b/internal/discord/bot.go @@ -0,0 +1,178 @@ +package discord + +import ( + "context" + "database/sql" + "errors" + "fmt" + "log" + "os" + "strings" + "time" + + "github.com/bwmarrin/discordgo" + + "plumber/internal/events" + "plumber/internal/store" +) + +const discordTimeout = 15 * time.Second + +// Bot posts site events to a Discord channel and owns post-to-message links. +type Bot struct { + channelID string + links store.DiscordLinkStore + api API +} + +// New constructs an outbound subscriber. Tests inject a fake API. +func New(channelID string, links store.DiscordLinkStore, api API) *Bot { + return &Bot{channelID: strings.TrimSpace(channelID), links: links, api: api} +} + +// 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) { + token := strings.TrimSpace(os.Getenv("DISCORD_BOT_TOKEN")) + channelID := strings.TrimSpace(os.Getenv("DISCORD_CHANNEL_ID")) + if token == "" && channelID == "" { + return nil, nil + } + if token == "" { + return nil, fmt.Errorf("DISCORD_BOT_TOKEN is required when DISCORD_CHANNEL_ID is set") + } + if channelID == "" { + return nil, fmt.Errorf("DISCORD_CHANNEL_ID is required when DISCORD_BOT_TOKEN is set") + } + if links == nil { + return nil, fmt.Errorf("discord links store is required") + } + session, err := discordgo.New("Bot " + token) + if err != nil { + return nil, err + } + bot := New(channelID, links, &sessionAPI{session: session}) + if bus != nil { + bus.Subscribe(bot.Handle) + } + log.Printf("discord: outbound subscriber enabled") + return bot, nil +} + +// Close releases the Discord session. +func (b *Bot) Close() error { + if b == nil || b.api == nil { + return nil + } + return b.api.Close() +} + +// Handle processes one site event. Failures are logged and do not fail the request. +func (b *Bot) Handle(_ context.Context, ev any) { + if b == nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), discordTimeout) + defer cancel() + switch e := ev.(type) { + case events.PostCreated: + b.onCreated(ctx, e.PostEvent) + case events.PostUpdated: + b.onUpdated(ctx, e.PostEvent) + } +} + +func (b *Bot) onCreated(ctx context.Context, ev events.PostEvent) { + if isRoot(ev) { + b.createRoot(ctx, ev) + return + } + b.createReply(ctx, ev) +} + +func (b *Bot) onUpdated(ctx context.Context, ev events.PostEvent) { + link, err := b.links.GetByPostID(ctx, ev.PostID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + b.onCreated(ctx, ev) + return + } + log.Printf("discord: load link %s: %v", ev.PostID, err) + return + } + channelID, err := b.editChannel(ctx, ev, link) + if err != nil { + log.Printf("discord: edit channel %s: %v", ev.PostID, err) + return + } + if err := b.api.Edit(ctx, channelID, link.MessageID, formatMessage(ev)); err != nil { + log.Printf("discord: edit %s: %v", ev.PostID, err) + return + } + log.Printf("discord: edited %s", ev.PostID) +} + +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) + if err != nil { + log.Printf("discord: start thread %s: %v", ev.PostID, err) + return + } + if err := b.links.Upsert(ctx, store.DiscordLink{ + PostID: ev.PostID, + MessageID: messageID, + ThreadID: threadID, + }); err != nil { + log.Printf("discord: save root link %s: %v", ev.PostID, err) + return + } + log.Printf("discord: posted root %s", ev.PostID) +} + +func (b *Bot) createReply(ctx context.Context, ev events.PostEvent) { + root, err := b.links.GetByPostID(ctx, ev.RootID) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + log.Printf("discord: skip reply %s: no root thread", ev.PostID) + return + } + log.Printf("discord: load root link %s: %v", ev.RootID, err) + return + } + if strings.TrimSpace(root.ThreadID) == "" { + log.Printf("discord: skip reply %s: no root thread", ev.PostID) + return + } + messageID, err := b.api.SendToThread(ctx, root.ThreadID, formatMessage(ev)) + if err != nil { + log.Printf("discord: send reply %s: %v", ev.PostID, err) + return + } + if err := b.links.Upsert(ctx, store.DiscordLink{ + PostID: ev.PostID, + MessageID: messageID, + }); err != nil { + log.Printf("discord: save reply link %s: %v", ev.PostID, err) + return + } + log.Printf("discord: posted reply %s", ev.PostID) +} + +func (b *Bot) editChannel(ctx context.Context, ev events.PostEvent, link *store.DiscordLink) (string, error) { + if strings.TrimSpace(link.ThreadID) != "" { + return b.channelID, nil + } + root, err := b.links.GetByPostID(ctx, ev.RootID) + if err != nil { + return "", err + } + if strings.TrimSpace(root.ThreadID) == "" { + return "", fmt.Errorf("root %s has no thread", ev.RootID) + } + return root.ThreadID, nil +} diff --git a/internal/discord/bot_test.go b/internal/discord/bot_test.go new file mode 100644 index 0000000..2c5e178 --- /dev/null +++ b/internal/discord/bot_test.go @@ -0,0 +1,247 @@ +package discord + +import ( + "context" + "sync" + "testing" + + "strconv" + + "plumber/internal/events" + "plumber/internal/store" +) + +type recordedSend struct { + Kind string + ChannelID string + Name string + Msg Message +} + +type fakeAPI struct { + mu sync.Mutex + sends []recordedSend + edits []recordedSend + next int + failSend error +} + +func (f *fakeAPI) SendToChannel(_ context.Context, channelID string, msg Message) (string, error) { + return f.record("channel", channelID, "", msg) +} + +func (f *fakeAPI) StartThread(_ context.Context, channelID, messageID, name string) (string, error) { + f.mu.Lock() + defer f.mu.Unlock() + f.next++ + f.sends = append(f.sends, recordedSend{ + Kind: "thread", + ChannelID: channelID, + Name: name, + Msg: Message{ThreadName: name, URL: messageID}, + }) + return "thread-" + messageID, nil +} + +func (f *fakeAPI) SendToThread(_ context.Context, threadID string, msg Message) (string, error) { + return f.record("thread-msg", threadID, "", msg) +} + +func (f *fakeAPI) Edit(_ context.Context, channelID, messageID string, msg Message) error { + f.mu.Lock() + defer f.mu.Unlock() + f.edits = append(f.edits, recordedSend{ + Kind: "edit", + ChannelID: channelID, + Name: messageID, + Msg: msg, + }) + return nil +} + +func (f *fakeAPI) Close() error { return nil } + +func (f *fakeAPI) record(kind, channelID, name string, msg Message) (string, error) { + f.mu.Lock() + defer f.mu.Unlock() + if f.failSend != nil { + return "", f.failSend + } + f.next++ + id := "msg-" + strconv.Itoa(f.next) + f.sends = append(f.sends, recordedSend{Kind: kind, ChannelID: channelID, Name: name, Msg: msg}) + return id, nil +} + +func TestOutboundRootReplyAndEdit(t *testing.T) { + t.Parallel() + + links := newMemoryLinks() + api := &fakeAPI{} + bot := New("channel-1", links, api) + ctx := context.Background() + + root := events.PostEvent{ + PostID: "root-1", + RootID: "root-1", + Title: "Leaky sink", + Body: "Water under the cabinet.", + City: "Oakland", + AuthorName: "sam", + Permalink: "https://www.askaplumberfirst.com/questions/root-1#post-root-1", + Images: []events.Image{{URL: "https://cdn.example/a.jpg"}, {URL: "https://cdn.example/b.jpg"}}, + } + bot.Handle(ctx, events.PostCreated{PostEvent: root}) + + if len(api.sends) != 2 || api.sends[0].Kind != "channel" || api.sends[1].Kind != "thread" { + t.Fatalf("root sends = %+v", api.sends) + } + if api.sends[0].ChannelID != "channel-1" || api.sends[1].Name != "Leaky sink" { + t.Fatalf("root routing = %+v", api.sends) + } + if got := api.sends[0].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" { + t.Fatalf("root link = %+v, %v", link, err) + } + + reply := events.PostEvent{ + PostID: "reply-1", + RootID: "root-1", + ParentID: "root-1", + Body: "Replace the cartridge.", + AuthorName: "plumber", + 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" { + t.Fatalf("reply sends = %+v", api.sends) + } + replyLink, err := links.GetByPostID(ctx, "reply-1") + if err != nil || replyLink.MessageID != "msg-3" || replyLink.ThreadID != "" { + t.Fatalf("reply link = %+v, %v", replyLink, err) + } + + 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" { + t.Fatalf("root edit = %+v", api.edits) + } + if api.edits[0].Msg.Description != "Updated leak." { + t.Fatalf("root edit body = %+v", api.edits[0].Msg) + } + + 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" { + t.Fatalf("reply edit = %+v", api.edits) + } +} + +func TestOutboundSkipsReplyWithoutRootLink(t *testing.T) { + t.Parallel() + + api := &fakeAPI{} + bot := New("channel-1", newMemoryLinks(), api) + bot.Handle(context.Background(), events.PostCreated{PostEvent: events.PostEvent{ + PostID: "reply-1", + RootID: "missing", + ParentID: "missing", + Body: "Orphan reply", + }}) + if len(api.sends) != 0 { + t.Fatalf("unexpected sends %+v", api.sends) + } +} + +func TestOutboundUpdateWithoutLinkCreates(t *testing.T) { + t.Parallel() + + links := newMemoryLinks() + api := &fakeAPI{} + bot := New("channel-1", links, api) + bot.Handle(context.Background(), events.PostUpdated{PostEvent: events.PostEvent{ + PostID: "root-2", + RootID: "root-2", + Title: "Late question", + Body: "Created while Discord was down.", + }}) + link, err := links.GetByPostID(context.Background(), "root-2") + if err != nil || link.ThreadID == "" || len(api.sends) != 2 { + t.Fatalf("late create link=%+v sends=%+v err=%v", link, api.sends, err) + } +} + +func TestFormatMessage(t *testing.T) { + t.Parallel() + + got := formatMessage(events.PostEvent{ + Title: "Leaky sink", + Body: "It drips.", + City: "Oakland", + AuthorName: "sam", + Permalink: "https://example.com/q", + Images: []events.Image{{URL: "https://cdn.example/a.jpg", Description: "ignored"}}, + }) + if got.Title != "Leaky sink" || + got.Description != "It drips." || + got.City != "Oakland" || + got.Author != "sam" || + got.URL != "https://example.com/q" || + got.ThreadName != "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" { + t.Fatalf("reply format = %+v", reply) + } +} + +func TestFromEnvDisabled(t *testing.T) { + t.Setenv("DISCORD_BOT_TOKEN", "") + t.Setenv("DISCORD_CHANNEL_ID", "") + bot, err := FromEnv(newMemoryLinks(), nil) + if err != nil || bot != nil { + t.Fatalf("disabled FromEnv = (%v, %v)", bot, err) + } +} + +func TestFromEnvRequiresBoth(t *testing.T) { + t.Setenv("DISCORD_BOT_TOKEN", "token") + t.Setenv("DISCORD_CHANNEL_ID", "") + if _, err := FromEnv(newMemoryLinks(), 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 { + t.Fatal("expected error when token is missing") + } +} + +func TestMemoryLinkUpsertKeepsThread(t *testing.T) { + t.Parallel() + + links := newMemoryLinks() + ctx := context.Background() + if err := links.Upsert(ctx, store.DiscordLink{PostID: "p", MessageID: "m1", ThreadID: "t1"}); err != nil { + t.Fatal(err) + } + if err := links.Upsert(ctx, store.DiscordLink{PostID: "p", MessageID: "m2"}); err != nil { + t.Fatal(err) + } + got, err := links.GetByPostID(ctx, "p") + if err != nil || got.MessageID != "m2" || got.ThreadID != "t1" { + t.Fatalf("upsert keep thread = %+v, %v", got, err) + } + if _, err := links.GetByMessageID(ctx, "m2"); err != nil { + t.Fatal(err) + } + if _, err := links.GetRootByThreadID(ctx, "t1"); err != nil { + t.Fatal(err) + } +} diff --git a/internal/discord/format.go b/internal/discord/format.go new file mode 100644 index 0000000..f552a19 --- /dev/null +++ b/internal/discord/format.go @@ -0,0 +1,75 @@ +package discord + +import ( + "strings" + + "plumber/internal/events" +) + +const ( + embedTitleLimit = 256 + embedDescriptionLimit = 4096 + threadNameLimit = 100 + embedColor = 0xe96a26 +) + +// Message is a Discord-ready snapshot of a site post event. +type Message struct { + Title string + URL string + Description string + City string + Author string + ImageURLs []string + ThreadName string +} + +func formatMessage(ev events.PostEvent) Message { + title := strings.TrimSpace(ev.Title) + if title == "" { + title = "Reply" + } + author := strings.TrimSpace(ev.AuthorName) + if author == "" { + author = "Someone" + } + msg := Message{ + Title: truncateRunes(title, embedTitleLimit), + URL: strings.TrimSpace(ev.Permalink), + Description: truncateRunes(strings.TrimSpace(ev.Body), embedDescriptionLimit), + City: strings.TrimSpace(ev.City), + Author: author, + ThreadName: threadName(ev.Title), + } + for _, img := range ev.Images { + url := strings.TrimSpace(img.URL) + if url == "" { + continue + } + msg.ImageURLs = append(msg.ImageURLs, url) + } + return msg +} + +func threadName(title string) string { + title = strings.TrimSpace(title) + if title == "" { + return "Question" + } + return truncateRunes(title, threadNameLimit) +} + +func truncateRunes(s string, max int) string { + if max <= 0 { + return "" + } + runes := []rune(s) + if len(runes) <= max { + return s + } + return string(runes[:max]) +} + +func isRoot(ev events.PostEvent) bool { + return strings.TrimSpace(ev.ParentID) == "" +} diff --git a/internal/discord/memory_links.go b/internal/discord/memory_links.go new file mode 100644 index 0000000..8555f3d --- /dev/null +++ b/internal/discord/memory_links.go @@ -0,0 +1,88 @@ +package discord + +import ( + "context" + "database/sql" + "strings" + "sync" + "time" + + "plumber/internal/store" +) + +// memoryLinks is an in-process DiscordLinkStore for tests. +type memoryLinks struct { + mu sync.Mutex + byPost map[string]store.DiscordLink + byMessage map[string]string + byThread map[string]string +} + +func newMemoryLinks() *memoryLinks { + return &memoryLinks{ + byPost: map[string]store.DiscordLink{}, + byMessage: map[string]string{}, + byThread: map[string]string{}, + } +} + +func (m *memoryLinks) GetByPostID(_ context.Context, postID string) (*store.DiscordLink, error) { + m.mu.Lock() + defer m.mu.Unlock() + link, ok := m.byPost[strings.TrimSpace(postID)] + if !ok { + return nil, sql.ErrNoRows + } + cp := link + return &cp, nil +} + +func (m *memoryLinks) GetByMessageID(_ context.Context, messageID string) (*store.DiscordLink, error) { + m.mu.Lock() + defer m.mu.Unlock() + postID, ok := m.byMessage[strings.TrimSpace(messageID)] + if !ok { + return nil, sql.ErrNoRows + } + link := m.byPost[postID] + cp := link + return &cp, nil +} + +func (m *memoryLinks) GetRootByThreadID(_ context.Context, threadID string) (*store.DiscordLink, error) { + m.mu.Lock() + defer m.mu.Unlock() + postID, ok := m.byThread[strings.TrimSpace(threadID)] + if !ok { + return nil, sql.ErrNoRows + } + link := m.byPost[postID] + cp := link + return &cp, nil +} + +func (m *memoryLinks) Upsert(_ context.Context, link store.DiscordLink) error { + m.mu.Lock() + defer m.mu.Unlock() + link.PostID = strings.TrimSpace(link.PostID) + link.MessageID = strings.TrimSpace(link.MessageID) + link.ThreadID = strings.TrimSpace(link.ThreadID) + if link.CreatedAt == "" { + link.CreatedAt = time.Now().UTC().Format(time.RFC3339Nano) + } + if prev, ok := m.byPost[link.PostID]; ok { + delete(m.byMessage, prev.MessageID) + if prev.ThreadID != "" { + delete(m.byThread, prev.ThreadID) + } + if link.ThreadID == "" { + link.ThreadID = prev.ThreadID + } + } + m.byPost[link.PostID] = link + m.byMessage[link.MessageID] = link.PostID + if link.ThreadID != "" { + m.byThread[link.ThreadID] = link.PostID + } + return nil +} diff --git a/internal/store/discord_link.go b/internal/store/discord_link.go new file mode 100644 index 0000000..ee68521 --- /dev/null +++ b/internal/store/discord_link.go @@ -0,0 +1,105 @@ +package store + +import ( + "context" + "database/sql" + "fmt" + "strings" + "time" + + "plumber/internal/store/sqlc" +) + +// DiscordLink is the bot-owned mapping from a site post to a Discord message. +type DiscordLink struct { + PostID string + MessageID string + ThreadID string + CreatedAt string +} + +// DiscordLinkStore is the mapping table used only by the Discord subscriber. +// It is not part of Store. +type DiscordLinkStore interface { + GetByPostID(ctx context.Context, postID string) (*DiscordLink, error) + GetByMessageID(ctx context.Context, messageID string) (*DiscordLink, error) + GetRootByThreadID(ctx context.Context, threadID string) (*DiscordLink, error) + Upsert(ctx context.Context, link DiscordLink) error +} + +// DiscordLinks implements DiscordLinkStore against Postgres. +type DiscordLinks struct { + db *sql.DB +} + +// NewDiscordLinks wraps db. It is independent of Store. +func NewDiscordLinks(db *sql.DB) *DiscordLinks { + return &DiscordLinks{db: db} +} + +// GetByPostID returns the link for a site post. +func (d *DiscordLinks) GetByPostID(ctx context.Context, postID string) (*DiscordLink, error) { + if d == nil || d.db == nil { + return nil, fmt.Errorf("discord links: no database") + } + row, err := sqlc.New(d.db).GetDiscordPostLinkByPostID(ctx, strings.TrimSpace(postID)) + if err != nil { + return nil, err + } + return discordLinkFromRow(row), nil +} + +// GetByMessageID returns the link for a Discord message. +func (d *DiscordLinks) GetByMessageID(ctx context.Context, messageID string) (*DiscordLink, error) { + if d == nil || d.db == nil { + return nil, fmt.Errorf("discord links: no database") + } + row, err := sqlc.New(d.db).GetDiscordPostLinkByMessageID(ctx, strings.TrimSpace(messageID)) + if err != nil { + return nil, err + } + return discordLinkFromRow(row), nil +} + +// GetRootByThreadID returns the root link for a Discord thread. +func (d *DiscordLinks) GetRootByThreadID(ctx context.Context, threadID string) (*DiscordLink, error) { + if d == nil || d.db == nil { + return nil, fmt.Errorf("discord links: no database") + } + row, err := sqlc.New(d.db).GetDiscordPostLinkByThreadID(ctx, strings.TrimSpace(threadID)) + if err != nil { + return nil, err + } + return discordLinkFromRow(row), nil +} + +// Upsert inserts or replaces the Discord IDs for a post. +func (d *DiscordLinks) Upsert(ctx context.Context, link DiscordLink) error { + if d == nil || d.db == nil { + return fmt.Errorf("discord links: no database") + } + link.PostID = strings.TrimSpace(link.PostID) + link.MessageID = strings.TrimSpace(link.MessageID) + link.ThreadID = strings.TrimSpace(link.ThreadID) + if link.PostID == "" || link.MessageID == "" { + return fmt.Errorf("discord links: post and message ids are required") + } + if link.CreatedAt == "" { + link.CreatedAt = time.Now().UTC().Format(time.RFC3339Nano) + } + return sqlc.New(d.db).UpsertDiscordPostLink(ctx, sqlc.UpsertDiscordPostLinkParams{ + PostID: link.PostID, + DiscordMessageID: link.MessageID, + DiscordThreadID: link.ThreadID, + CreatedAt: link.CreatedAt, + }) +} + +func discordLinkFromRow(row sqlc.DiscordPostLink) *DiscordLink { + return &DiscordLink{ + PostID: row.PostID, + MessageID: row.DiscordMessageID, + ThreadID: row.DiscordThreadID, + CreatedAt: row.CreatedAt, + } +} diff --git a/internal/store/migrate.go b/internal/store/migrate.go index 4d12c8a..28638e9 100644 --- a/internal/store/migrate.go +++ b/internal/store/migrate.go @@ -139,6 +139,25 @@ CREATE TABLE IF NOT EXISTS post_images ( return nil } +func migrateDiscordPostLinks(ctx context.Context, exec execContext) error { + if _, err := exec.ExecContext(ctx, ` +CREATE TABLE IF NOT EXISTS discord_post_links ( + post_id TEXT PRIMARY KEY REFERENCES posts(id) ON DELETE CASCADE, + discord_message_id TEXT NOT NULL UNIQUE, + discord_thread_id TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL +)`); err != nil { + return fmt.Errorf("create discord_post_links: %w", err) + } + if _, err := exec.ExecContext(ctx, ` +CREATE UNIQUE INDEX IF NOT EXISTS discord_post_links_thread_uidx + ON discord_post_links (discord_thread_id) + WHERE discord_thread_id <> ''`); err != nil { + return fmt.Errorf("discord_post_links_thread_uidx: %w", err) + } + return nil +} + func migratePostDate(ctx context.Context, exec execContext) error { steps := []struct { name string @@ -331,6 +350,7 @@ CREATE TABLE IF NOT EXISTS schema_migrations ( {"008_post_author_index", migratePostAuthorIndex}, {"009_drop_legacy_post_tables", migrateDropLegacyPostTables}, {"010_post_images", migratePostImages}, + {"011_discord_post_links", migrateDiscordPostLinks}, } for _, m := range migrations { if applied[m.version] { diff --git a/internal/store/migrate_posts_test.go b/internal/store/migrate_posts_test.go index 5db4972..d2dba7c 100644 --- a/internal/store/migrate_posts_test.go +++ b/internal/store/migrate_posts_test.go @@ -81,6 +81,12 @@ CREATE TABLE users ( if err := migratePostImages(ctx, conn); err != nil { t.Fatalf("post images migration is not idempotent: %v", err) } + if err := migrateDiscordPostLinks(ctx, conn); err != nil { + t.Fatal(err) + } + if err := migrateDiscordPostLinks(ctx, conn); err != nil { + t.Fatalf("discord post links migration is not idempotent: %v", err) + } if _, err := conn.ExecContext(ctx, ` INSERT INTO users (id, name, role) VALUES ('homeowner', 'Home Owner', 'user'), ('plumber', 'The Plumber', 'admin'); @@ -143,6 +149,42 @@ VALUES ('homeowner', 'root-1', 1);`); err != nil { }); err == nil { t.Fatal("fifth image position unexpectedly succeeded") } + if err := imageQueries.UpsertDiscordPostLink(ctx, sqlc.UpsertDiscordPostLinkParams{ + PostID: "root-1", + DiscordMessageID: "msg-root", + DiscordThreadID: "thread-root", + CreatedAt: "2026-08-26T08:00:00Z", + }); err != nil { + t.Fatal(err) + } + if err := imageQueries.UpsertDiscordPostLink(ctx, sqlc.UpsertDiscordPostLinkParams{ + PostID: "reply-1", + DiscordMessageID: "msg-reply", + DiscordThreadID: "", + CreatedAt: "2026-08-26T09:00:00Z", + }); err != nil { + t.Fatal(err) + } + rootLink, err := imageQueries.GetDiscordPostLinkByPostID(ctx, "root-1") + if err != nil || rootLink.DiscordMessageID != "msg-root" || rootLink.DiscordThreadID != "thread-root" { + t.Fatalf("root discord link = %+v, %v", rootLink, err) + } + threadLink, err := imageQueries.GetDiscordPostLinkByThreadID(ctx, "thread-root") + if err != nil || threadLink.PostID != "root-1" { + t.Fatalf("thread discord link = %+v, %v", threadLink, err) + } + if err := imageQueries.UpsertDiscordPostLink(ctx, sqlc.UpsertDiscordPostLinkParams{ + PostID: "reply-1", + DiscordMessageID: "msg-reply-2", + DiscordThreadID: "", + CreatedAt: "2026-08-26T09:01:00Z", + }); err != nil { + t.Fatal(err) + } + replyLink, err := imageQueries.GetDiscordPostLinkByPostID(ctx, "reply-1") + if err != nil || replyLink.DiscordMessageID != "msg-reply-2" || replyLink.DiscordThreadID != "" { + t.Fatalf("reply upsert = %+v, %v", replyLink, err) + } var postVoteIndexCount int if err := conn.QueryRowContext(ctx, ` SELECT count(*) diff --git a/internal/store/sqlc/discord_links.sql.go b/internal/store/sqlc/discord_links.sql.go new file mode 100644 index 0000000..79d1680 --- /dev/null +++ b/internal/store/sqlc/discord_links.sql.go @@ -0,0 +1,100 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.1 +// source: discord_links.sql + +package sqlc + +import ( + "context" +) + +const getDiscordPostLinkByMessageID = `-- name: GetDiscordPostLinkByMessageID :one +SELECT post_id, discord_message_id, discord_thread_id, created_at +FROM discord_post_links +WHERE discord_message_id = $1 +` + +func (q *Queries) GetDiscordPostLinkByMessageID(ctx context.Context, discordMessageID string) (DiscordPostLink, error) { + row := q.db.QueryRowContext(ctx, getDiscordPostLinkByMessageID, discordMessageID) + var i DiscordPostLink + err := row.Scan( + &i.PostID, + &i.DiscordMessageID, + &i.DiscordThreadID, + &i.CreatedAt, + ) + return i, err +} + +const getDiscordPostLinkByPostID = `-- name: GetDiscordPostLinkByPostID :one +SELECT post_id, discord_message_id, discord_thread_id, created_at +FROM discord_post_links +WHERE post_id = $1 +` + +func (q *Queries) GetDiscordPostLinkByPostID(ctx context.Context, postID string) (DiscordPostLink, error) { + row := q.db.QueryRowContext(ctx, getDiscordPostLinkByPostID, postID) + var i DiscordPostLink + err := row.Scan( + &i.PostID, + &i.DiscordMessageID, + &i.DiscordThreadID, + &i.CreatedAt, + ) + return i, err +} + +const getDiscordPostLinkByThreadID = `-- name: GetDiscordPostLinkByThreadID :one +SELECT post_id, discord_message_id, discord_thread_id, created_at +FROM discord_post_links +WHERE discord_thread_id = $1 + AND discord_thread_id <> '' +` + +func (q *Queries) GetDiscordPostLinkByThreadID(ctx context.Context, discordThreadID string) (DiscordPostLink, error) { + row := q.db.QueryRowContext(ctx, getDiscordPostLinkByThreadID, discordThreadID) + var i DiscordPostLink + err := row.Scan( + &i.PostID, + &i.DiscordMessageID, + &i.DiscordThreadID, + &i.CreatedAt, + ) + return i, err +} + +const upsertDiscordPostLink = `-- name: UpsertDiscordPostLink :exec +INSERT INTO discord_post_links ( + post_id, discord_message_id, discord_thread_id, created_at +) +VALUES ( + $1, + $2, + $3, + $4 +) +ON CONFLICT (post_id) DO UPDATE SET + discord_message_id = EXCLUDED.discord_message_id, + discord_thread_id = CASE + WHEN EXCLUDED.discord_thread_id <> '' THEN EXCLUDED.discord_thread_id + ELSE discord_post_links.discord_thread_id + END +` + +type UpsertDiscordPostLinkParams struct { + PostID string + DiscordMessageID string + DiscordThreadID string + CreatedAt string +} + +func (q *Queries) UpsertDiscordPostLink(ctx context.Context, arg UpsertDiscordPostLinkParams) error { + _, err := q.db.ExecContext(ctx, upsertDiscordPostLink, + arg.PostID, + arg.DiscordMessageID, + arg.DiscordThreadID, + arg.CreatedAt, + ) + return err +} diff --git a/internal/store/sqlc/models.go b/internal/store/sqlc/models.go index a9584f4..560d0ce 100644 --- a/internal/store/sqlc/models.go +++ b/internal/store/sqlc/models.go @@ -9,6 +9,13 @@ import ( "time" ) +type DiscordPostLink struct { + PostID string + DiscordMessageID string + DiscordThreadID string + CreatedAt string +} + type Post struct { ID string ParentID sql.NullString diff --git a/schema.sql b/schema.sql index a163643..4caac27 100644 --- a/schema.sql +++ b/schema.sql @@ -65,6 +65,17 @@ CREATE TABLE IF NOT EXISTS post_votes ( CREATE INDEX IF NOT EXISTS idx_post_votes_post_id ON post_votes(post_id); +CREATE TABLE IF NOT EXISTS discord_post_links ( + post_id TEXT PRIMARY KEY REFERENCES posts(id) ON DELETE CASCADE, + discord_message_id TEXT NOT NULL UNIQUE, + discord_thread_id TEXT NOT NULL DEFAULT '', + created_at TEXT NOT NULL +); + +CREATE UNIQUE INDEX IF NOT EXISTS discord_post_links_thread_uidx + ON discord_post_links (discord_thread_id) + WHERE discord_thread_id <> ''; + CREATE TABLE IF NOT EXISTS sessions ( token TEXT PRIMARY KEY, data BYTEA NOT NULL,