security: GHSA-xpxj-f2fm-rqch (HIGH) bound OAuth2 state map growth

Sweep expired OAuth2 state entries, cap the map at 10000 entries, and
remove stale state on failed callback validation.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
jamesread 2026-07-07 13:07:15 +01:00
parent 5bcbdb4085
commit 422044317c
2 changed files with 98 additions and 0 deletions

View File

@ -62,8 +62,14 @@ type oauth2State struct {
providerName string providerName string
Username string Username string
Usergroup string Usergroup string
createdAt time.Time
} }
const (
oauthStateMaxAge = 900 // matches olivetin-sid-oauth cookie MaxAge
oauthStateMaxEntries = 10000
)
func assignIfEmpty(target *string, value string) { func assignIfEmpty(target *string, value string) {
if *target == "" { if *target == "" {
*target = value *target = value
@ -129,6 +135,19 @@ func (h *OAuth2Handler) setOAuthCallbackCookie(w http.ResponseWriter, r *http.Re
http.SetCookie(w, cookie) http.SetCookie(w, cookie)
} }
func (h *OAuth2Handler) deleteOAuthStateLocked(state string) {
delete(h.registeredStates, state)
}
func (h *OAuth2Handler) sweepExpiredOAuthStatesLocked(now time.Time) {
cutoff := now.Add(-oauthStateMaxAge * time.Second)
for state, entry := range h.registeredStates {
if entry.createdAt.Before(cutoff) {
delete(h.registeredStates, state)
}
}
}
func (h *OAuth2Handler) HandleOAuthLogin(w http.ResponseWriter, r *http.Request) { func (h *OAuth2Handler) HandleOAuthLogin(w http.ResponseWriter, r *http.Request) {
state, err := randString(16) state, err := randString(16)
@ -147,10 +166,17 @@ func (h *OAuth2Handler) HandleOAuthLogin(w http.ResponseWriter, r *http.Request)
} }
h.mu.Lock() h.mu.Lock()
h.sweepExpiredOAuthStatesLocked(time.Now())
if len(h.registeredStates) >= oauthStateMaxEntries {
h.mu.Unlock()
http.Error(w, "OAuth login temporarily unavailable", http.StatusServiceUnavailable)
return
}
h.registeredStates[state] = &oauth2State{ h.registeredStates[state] = &oauth2State{
providerConfig: provider, providerConfig: provider,
providerName: providerName, providerName: providerName,
Username: "", Username: "",
createdAt: time.Now(),
} }
h.mu.Unlock() h.mu.Unlock()
@ -177,6 +203,9 @@ func (h *OAuth2Handler) checkOAuthCallbackCookie(w http.ResponseWriter, r *http.
if !h.validateStateMatch(r.URL.Query().Get("state"), state) { if !h.validateStateMatch(r.URL.Query().Get("state"), state) {
log.Errorf("State mismatch: %v != %v", r.URL.Query().Get("state"), state) log.Errorf("State mismatch: %v != %v", r.URL.Query().Get("state"), state)
h.mu.Lock()
h.deleteOAuthStateLocked(state)
h.mu.Unlock()
http.Error(w, "State mismatch", http.StatusBadRequest) http.Error(w, "State mismatch", http.StatusBadRequest)
return nil, state, false return nil, state, false
} }
@ -186,6 +215,9 @@ func (h *OAuth2Handler) checkOAuthCallbackCookie(w http.ResponseWriter, r *http.
h.mu.RUnlock() h.mu.RUnlock()
if !ok { if !ok {
log.Errorf("State not found in server: %v", state) log.Errorf("State not found in server: %v", state)
h.mu.Lock()
h.deleteOAuthStateLocked(state)
h.mu.Unlock()
http.Error(w, "State not found in server", http.StatusBadRequest) http.Error(w, "State not found in server", http.StatusBadRequest)
return nil, state, false return nil, state, false
} }

View File

@ -0,0 +1,66 @@
package otoauth2
import (
"net/http"
"net/http/httptest"
"strconv"
"testing"
"time"
config "github.com/OliveTin/OliveTin/internal/config"
"github.com/stretchr/testify/assert"
"golang.org/x/oauth2"
)
func TestSweepExpiredOAuthStatesLocked(t *testing.T) {
h := &OAuth2Handler{
registeredStates: make(map[string]*oauth2State),
}
h.registeredStates["fresh"] = &oauth2State{
providerName: "test",
createdAt: time.Now(),
}
h.registeredStates["stale"] = &oauth2State{
providerName: "test",
createdAt: time.Now().Add(-2 * oauthStateMaxAge * time.Second),
}
h.sweepExpiredOAuthStatesLocked(time.Now())
_, freshFound := h.registeredStates["fresh"]
_, staleFound := h.registeredStates["stale"]
assert.True(t, freshFound)
assert.False(t, staleFound)
}
func TestHandleOAuthLoginRejectsWhenStateMapFull(t *testing.T) {
cfg := config.DefaultConfig()
cfg.AuthOAuth2Providers = map[string]*config.OAuth2Provider{
"test": {
Name: "test",
ClientID: "id",
ClientSecret: "secret",
AuthUrl: "https://example.com/auth",
TokenUrl: "https://example.com/token",
},
}
h := NewOAuth2Handler(cfg)
h.registeredStates = make(map[string]*oauth2State, oauthStateMaxEntries)
for i := 0; i < oauthStateMaxEntries; i++ {
h.registeredStates[strconv.Itoa(i)] = &oauth2State{
providerConfig: &oauth2.Config{},
providerName: "test",
createdAt: time.Now(),
}
}
req := httptest.NewRequest(http.MethodGet, "/oauth/login?provider=test", nil)
rec := httptest.NewRecorder()
h.HandleOAuthLogin(rec, req)
assert.Equal(t, http.StatusServiceUnavailable, rec.Code)
assert.Equal(t, oauthStateMaxEntries, len(h.registeredStates))
}