|
|
|
@@ -9,6 +9,7 @@ import (
|
|
|
|
|
"log"
|
|
|
|
|
"net/http"
|
|
|
|
|
"net/url"
|
|
|
|
|
"slices"
|
|
|
|
|
"strconv"
|
|
|
|
|
"strings"
|
|
|
|
|
"sync"
|
|
|
|
@@ -43,9 +44,9 @@ type DiscordConfig struct {
|
|
|
|
|
// RequiredRole, when non-empty, is a role ID a member must hold on top of
|
|
|
|
|
// guild membership. Empty by default: membership alone suffices.
|
|
|
|
|
RequiredRole string
|
|
|
|
|
// APIBBase is the Discord API root; configurable so tests run the whole
|
|
|
|
|
// APIBase is the Discord API root; configurable so tests run the whole
|
|
|
|
|
// flow against a local stub.
|
|
|
|
|
APIBBase string
|
|
|
|
|
APIBase string
|
|
|
|
|
// RedirectURI is the full public URL of the callback — Discord requires
|
|
|
|
|
// the exact string, so it is configured, never derived from headers.
|
|
|
|
|
RedirectURI string
|
|
|
|
@@ -59,7 +60,6 @@ type DiscordConfig struct {
|
|
|
|
|
type oauthStates struct {
|
|
|
|
|
mu sync.Mutex
|
|
|
|
|
expiry map[string]time.Time
|
|
|
|
|
order []string // FIFO for eviction when the map is full
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func newOAuthStates() *oauthStates {
|
|
|
|
@@ -75,18 +75,19 @@ func (s *oauthStates) put(state string, expires time.Time) {
|
|
|
|
|
delete(s.expiry, k)
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
// Evict the state closest to expiring when full, so a flood of starts
|
|
|
|
|
// cannot grow memory; at worst it invalidates an in-flight sign-in.
|
|
|
|
|
if len(s.expiry) >= maxStates {
|
|
|
|
|
// Evict from the front until under the cap. The front may already
|
|
|
|
|
// have been consumed by take (which removes from the map, not the
|
|
|
|
|
// order slice), so pop until the count actually drops.
|
|
|
|
|
for len(s.expiry) >= maxStates && len(s.order) > 0 {
|
|
|
|
|
oldest := s.order[0]
|
|
|
|
|
s.order = s.order[1:]
|
|
|
|
|
delete(s.expiry, oldest)
|
|
|
|
|
var oldest string
|
|
|
|
|
var oldestAt time.Time
|
|
|
|
|
for k, at := range s.expiry {
|
|
|
|
|
if oldest == "" || at.Before(oldestAt) {
|
|
|
|
|
oldest, oldestAt = k, at
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
delete(s.expiry, oldest)
|
|
|
|
|
}
|
|
|
|
|
s.expiry[state] = expires
|
|
|
|
|
s.order = append(s.order, state)
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
// take validates and consumes a state in one step: a state works exactly
|
|
|
|
@@ -107,7 +108,7 @@ func (s *oauthStates) take(state string) bool {
|
|
|
|
|
func (h *Handler) discordStart(w http.ResponseWriter, r *http.Request) {
|
|
|
|
|
state := session.NewID()
|
|
|
|
|
h.states.put(state, time.Now().Add(oauthStateTTL))
|
|
|
|
|
u := h.discord.APIBBase + "/oauth2/authorize?" + url.Values{
|
|
|
|
|
u := h.discord.APIBase + "/oauth2/authorize?" + url.Values{
|
|
|
|
|
"client_id": {h.discord.ClientID},
|
|
|
|
|
"redirect_uri": {h.discord.RedirectURI},
|
|
|
|
|
"response_type": {"code"},
|
|
|
|
@@ -184,7 +185,7 @@ func (h *Handler) discordCallback(w http.ResponseWriter, r *http.Request) {
|
|
|
|
|
// The refusal is the same for a non-member and a member without the
|
|
|
|
|
// required role, and it names neither the guild nor its id: an outsider
|
|
|
|
|
// cannot tell whether the guild exists, let alone which one gates.
|
|
|
|
|
if !isMember || (h.discord.RequiredRole != "" && !contains(member.Roles, h.discord.RequiredRole)) {
|
|
|
|
|
if !isMember || (h.discord.RequiredRole != "" && !slices.Contains(member.Roles, h.discord.RequiredRole)) {
|
|
|
|
|
h.limiter.Fail(ip, time.Now())
|
|
|
|
|
h.renderLogin(w, http.StatusForbidden,
|
|
|
|
|
"This Discord account is not a member of this community.")
|
|
|
|
@@ -214,7 +215,7 @@ func (h *Handler) exchangeToken(ctx context.Context, code string) (discordToken,
|
|
|
|
|
"redirect_uri": {h.discord.RedirectURI},
|
|
|
|
|
}
|
|
|
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
|
|
|
|
|
h.discord.APIBBase+"/oauth2/token", strings.NewReader(form.Encode()))
|
|
|
|
|
h.discord.APIBase+"/oauth2/token", strings.NewReader(form.Encode()))
|
|
|
|
|
if err != nil {
|
|
|
|
|
return discordToken{}, err
|
|
|
|
|
}
|
|
|
|
@@ -241,7 +242,7 @@ func (h *Handler) exchangeToken(ctx context.Context, code string) (discordToken,
|
|
|
|
|
// discordUserID fetches the signed-in user's id via the identify scope.
|
|
|
|
|
func (h *Handler) discordUserID(ctx context.Context, accessToken string) (string, error) {
|
|
|
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
|
|
|
|
|
h.discord.APIBBase+"/users/@me", nil)
|
|
|
|
|
h.discord.APIBase+"/users/@me", nil)
|
|
|
|
|
if err != nil {
|
|
|
|
|
return "", err
|
|
|
|
|
}
|
|
|
|
@@ -276,7 +277,7 @@ type discordMember struct {
|
|
|
|
|
// gate asks exactly the question it names. A 404 or 403 (not in the guild, or
|
|
|
|
|
// the token lacks the scope) is a non-member, not an error.
|
|
|
|
|
func (h *Handler) discordMember(ctx context.Context, accessToken, userID string) (discordMember, bool, error) {
|
|
|
|
|
u := h.discord.APIBBase + "/guilds/" + url.PathEscape(h.discord.GuildID) +
|
|
|
|
|
u := h.discord.APIBase + "/guilds/" + url.PathEscape(h.discord.GuildID) +
|
|
|
|
|
"/members/" + url.PathEscape(userID)
|
|
|
|
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, u, nil)
|
|
|
|
|
if err != nil {
|
|
|
|
@@ -305,12 +306,3 @@ func (h *Handler) discordMember(ctx context.Context, accessToken, userID string)
|
|
|
|
|
type discordToken struct {
|
|
|
|
|
AccessToken string `json:"access_token"`
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
func contains(ss []string, want string) bool {
|
|
|
|
|
for _, s := range ss {
|
|
|
|
|
if s == want {
|
|
|
|
|
return true
|
|
|
|
|
}
|
|
|
|
|
}
|
|
|
|
|
return false
|
|
|
|
|
}
|
|
|
|
|