Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 37 additions & 10 deletions platform/discord/discord.go
Original file line number Diff line number Diff line change
Expand Up @@ -507,11 +507,19 @@ func (p *Platform) RegisterCommands(commands []core.BotCommandInfo) error {
return nil
}

registered, err := p.session.ApplicationCommandBulkOverwrite(p.appID, p.guildID, cmds)
p.mu.RLock()
session := p.session
appID := p.appID
p.mu.RUnlock()
if session == nil {
return fmt.Errorf("discord: session not connected")
}

registered, err := session.ApplicationCommandBulkOverwrite(appID, p.guildID, cmds)
if err != nil {
slog.Error("discord: failed to register slash commands — "+
"make sure the bot was invited with BOTH 'bot' AND 'applications.commands' OAuth2 scopes. "+
"Re-invite URL: https://discord.com/oauth2/authorize?client_id="+p.appID+
"Re-invite URL: https://discord.com/oauth2/authorize?client_id="+appID+
"&scope=bot+applications.commands&permissions=2147485696",
"error", err, "guild_id", p.guildID)
return err
Expand Down Expand Up @@ -576,8 +584,12 @@ func (p *Platform) buildSession() (*discordgo.Session, error) {
session.Identify.Intents = discordgo.IntentsGuilds | discordgo.IntentsGuildMessages | discordgo.IntentsDirectMessages | discordgo.IntentMessageContent

session.AddHandler(func(s *discordgo.Session, r *discordgo.Ready) {
// botID/appID are read from MessageCreate / GuildCreate / RegisterCommands
// which may run concurrently with this Ready callback; take the write lock.
p.mu.Lock()
p.botID = r.User.ID
p.appID = r.User.ID
p.mu.Unlock()
slog.Info("discord: connected", "bot", r.User.Username+"#"+r.User.Discriminator)
// Signal readiness before guild role lookups so RegisterCommands
// is not blocked by slow API calls when there are many guilds.
Expand All @@ -602,13 +614,22 @@ func (p *Platform) buildSession() (*discordgo.Session, error) {
})

session.AddHandler(func(s *discordgo.Session, m *discordgo.MessageCreate) {
// Snapshot botID/session under the read lock. The Ready handler
// writes botID concurrently with this handler, and the connect
// loop swaps p.session on reconnect; reading without the lock
// races on both fields under -race.
p.mu.RLock()
botID := p.botID
connSession := p.session
p.mu.RUnlock()

// Deduplicate: Discord gateway may deliver the same event twice
if !rememberDedupID(&p.seenMsgs, m.ID) {
slog.Debug("discord: ignoring duplicate message", "msg_id", m.ID)
return
}

if m.Author.Bot || m.Author.ID == p.botID {
if m.Author.Bot || m.Author.ID == botID {
return
}
if core.IsOldMessage(m.Timestamp) {
Expand All @@ -629,11 +650,11 @@ func (p *Platform) buildSession() (*discordgo.Session, error) {
botRoleID = p.botRoleIDForGuild(m.GuildID)
}
if m.GuildID != "" && !p.isGroupReplyAllGuild(m.GuildID) {
if !isDiscordBotMention(m, p.botID, botRoleID, p.respondToAtEveryoneAndHere) {
if !isDiscordBotMention(m, botID, botRoleID, p.respondToAtEveryoneAndHere) {
slog.Debug("discord: ignoring guild message without bot mention", "channel", m.ChannelID)
return
}
m.Content = stripDiscordMentionWithRole(m.Content, p.botID, botRoleID)
m.Content = stripDiscordMentionWithRole(m.Content, botID, botRoleID)
if m.MentionEveryone {
m.Content = stripEveryoneHere(m.Content)
}
Expand All @@ -651,7 +672,7 @@ func (p *Platform) buildSession() (*discordgo.Session, error) {
// (the historical, non-isolated behavior).
channelKey := ""
if p.threadIsolation && m.GuildID != "" {
threadSessionKey, threadCtx, parentChannelID, err := resolveThreadReplyContext(m, p.botID, sessionThreadOps{session: p.session})
threadSessionKey, threadCtx, parentChannelID, err := resolveThreadReplyContext(m, botID, sessionThreadOps{session: connSession})
if err != nil {
slog.Warn("discord: thread isolation setup failed, falling back", "message", m.ID, "channel", m.ChannelID, "error", err)
} else {
Expand Down Expand Up @@ -1395,10 +1416,16 @@ func (p *Platform) botRoleIDForGuild(guildID string) string {
}

func (p *Platform) cacheBotRoleIDForGuild(s *discordgo.Session, guildID string, guildRoles []*discordgo.Role) {
if s == nil || guildID == "" || p.botID == "" {
if s == nil || guildID == "" {
return
}
p.mu.RLock()
botID := p.botID
p.mu.RUnlock()
if botID == "" {
return
}
roleID, err := p.resolveBotRoleIDForGuild(s, guildID, guildRoles)
roleID, err := p.resolveBotRoleIDForGuild(s, guildID, botID, guildRoles)
if err != nil {
slog.Debug("discord: resolve bot managed role failed", "guild", guildID, "error", err)
return
Expand All @@ -1409,8 +1436,8 @@ func (p *Platform) cacheBotRoleIDForGuild(s *discordgo.Session, guildID string,
p.botRoleIDs.Store(guildID, roleID)
}

func (p *Platform) resolveBotRoleIDForGuild(s *discordgo.Session, guildID string, guildRoles []*discordgo.Role) (string, error) {
member, err := s.GuildMember(guildID, p.botID)
func (p *Platform) resolveBotRoleIDForGuild(s *discordgo.Session, guildID, botID string, guildRoles []*discordgo.Role) (string, error) {
member, err := s.GuildMember(guildID, botID)
if err != nil {
return "", fmt.Errorf("fetch bot member: %w", err)
}
Expand Down
78 changes: 78 additions & 0 deletions platform/discord/identity_race_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
package discord

import (
"errors"
"net/http"
"sync"
"testing"

"github.com/bwmarrin/discordgo"
)

type roundTripperFunc func(*http.Request) (*http.Response, error)

func (f roundTripperFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) }

// Regression: the Ready callback wrote p.botID/p.appID without holding p.mu
// while MessageCreate, GuildCreate, RegisterCommands, and cacheBotRoleIDForGuild
// read those fields from separate goroutines. discordgo dispatches each event
// in its own goroutine, so -race flagged the field accesses.
//
// The writer goroutine mirrors the fixed Ready handler (write under p.mu.Lock).
// The reader goroutine exercises cacheBotRoleIDForGuild, which must take
// p.mu.RLock before reading p.botID; -race fails loudly if that read is ever
// reverted to an unlocked access.
//
// Run with: go test -race ./platform/discord/ -run TestPlatform_IdentityFields_Race
func TestPlatform_IdentityFields_Race(t *testing.T) {
// Session whose HTTP client fails instantly so cacheBotRoleIDForGuild
// never blocks on a network round-trip even when it sees a non-empty botID.
session, err := discordgo.New("Bot test-token")
if err != nil {
t.Fatalf("discordgo.New: %v", err)
}
session.Client = &http.Client{Transport: roundTripperFunc(func(*http.Request) (*http.Response, error) {
return nil, errors.New("no network in race test")
})}

p := &Platform{
token: "test-token",
session: session,
readyCh: make(chan struct{}),
}
p.mu.Lock()
p.botID = "bot-seed"
p.appID = "app-seed"
p.mu.Unlock()

const goroutines = 6
const iters = 50
var wg sync.WaitGroup
wg.Add(goroutines * 2)

for i := 0; i < goroutines; i++ {
go func() {
defer wg.Done()
for j := 0; j < iters; j++ {
// Mirror the fixed Ready handler: identity writes under Lock.
p.mu.Lock()
p.botID = "bot-123"
p.appID = "app-456"
p.mu.Unlock()
}
}()
}

for i := 0; i < goroutines; i++ {
go func() {
defer wg.Done()
for j := 0; j < iters; j++ {
// Production code path; must take p.mu.RLock internally.
// If the RLock around p.botID is ever removed, -race fails.
p.cacheBotRoleIDForGuild(session, "guild-1", nil)
}
}()
}

wg.Wait()
}
Loading