Skip to content
Merged
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
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -583,7 +583,8 @@ Response codes:
tell an attacker something.
- `415 Unsupported Media Type` — non-JSON `Content-Type`.
- `429 Too Many Requests` — per-IP rate limit. The endpoint is unauthenticated,
so it shares the login rate limiter's per-IP budget.
so it is limited like login, but on a separate budget of 60 a minute: a
burst of refreshes from one address never blocks logins from it.
- `500 Internal Server Error` — the token could not be read or rotated, or
revoking the user's tokens after a replay failed.

Expand Down
13 changes: 7 additions & 6 deletions auth_refresh.go
Original file line number Diff line number Diff line change
Expand Up @@ -65,14 +65,15 @@ type RefreshStore interface {
// The route is unauthenticated by necessity — the access token is expired by
// definition at the moment a client needs this, so the refresh token is the
// credential. That makes it as exposed as login, and it gets login's
// defenses: a 1KB body cap, strict JSON decoding, and the same rate limiter.
// defenses: a 1KB body cap, strict JSON decoding, and a per-IP rate limit,
// though from its own RefreshRateLimiter rather than login's.
//
// Tokens are single-use. Each refresh consumes the presented token and
// returns its replacement, so a stolen token stops working as soon as the
// legitimate client refreshes. Presenting a token that is already spent
// revokes every refresh token its user holds, because the server cannot tell
// a client retrying a lost response from a thief replaying a stolen token.
func handleRefreshToken(store RefreshStore, secret []byte, ttls tokenTTLs, limiter *LoginRateLimiter, trustProxy bool) http.HandlerFunc {
func handleRefreshToken(store RefreshStore, secret []byte, ttls tokenTTLs, limiter *RefreshRateLimiter, trustProxy bool) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
mediaType, _, err := mime.ParseMediaType(r.Header.Get("Content-Type"))
if err != nil || !strings.EqualFold(mediaType, "application/json") {
Expand Down Expand Up @@ -106,10 +107,10 @@ func handleRefreshToken(store RefreshStore, secret []byte, ttls tokenTTLs, limit

// Only the IP dimension applies: the caller presents a token, not an
// email, and there is nothing to key a per-account window on until
// the lookup below. This shares the login limiter's per-IP budget by
// design — one address cannot buy itself extra attempts by spreading
// them across the two auth endpoints.
if limiter != nil && !limiter.AllowIP(ip) {
// the lookup below. The budget is refresh's own, not login's, so a
// depot whose drivers refresh together cannot lock itself out of
// logging in.
if limiter != nil && !limiter.Allow(ip) {
writeJSON(w, http.StatusTooManyRequests, map[string]string{"error": "too many attempts"})
return
}
Expand Down
53 changes: 49 additions & 4 deletions auth_refresh_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -468,15 +468,15 @@ func TestHandleRefresh_ConcurrentLossRevokes(t *testing.T) {

func TestHandleRefresh_RateLimited(t *testing.T) {
f, token := refreshStoreWithUser(t)
limiter := NewLoginRateLimiter()
limiter := NewRefreshRateLimiter()
defer limiter.Stop()

handler := handleRefreshToken(f, testSecret, testTTLs, limiter, false)

// Each refresh rotates, so present the token the previous call returned
// and stay on the success path until the limiter itself trips.
current := token
for range loginIPLimit {
for range refreshIPLimit {
w := postRefresh(handler, current)
require.Equal(t, http.StatusOK, w.Code)
current = decodeTokens(t, w).RefreshToken
Expand All @@ -492,18 +492,63 @@ func TestHandleRefresh_RateLimited(t *testing.T) {
func TestHandleRefresh_RateLimitedBeforeStore(t *testing.T) {
f, token := refreshStoreWithUser(t)
f.getErr = errors.New("the store must not be reached")
limiter := NewLoginRateLimiter()
limiter := NewRefreshRateLimiter()
defer limiter.Stop()

handler := handleRefreshToken(f, testSecret, testTTLs, limiter, false)
for range loginIPLimit {
for range refreshIPLimit {
require.Equal(t, http.StatusInternalServerError, postRefresh(handler, token).Code)
}

w := postRefresh(handler, token)
assert.Equal(t, http.StatusTooManyRequests, w.Code)
}

// loginStubStore lets a login through the real mux succeed for one user.
// Refresh tokens are all unknown, as noopStore reports them.
type loginStubStore struct {
noopStore
user *User
}

func (s *loginStubStore) GetUserByEmail(_ context.Context, email string) (*User, error) {
if email != s.user.Email {
return nil, ErrUserNotFound
}
return s.user, nil
}

// TestRefresh_DoesNotSpendLoginBudget drives the real mux, so it pins which
// limiter the refresh route is wired to: an address that has used up its
// refresh allowance must still be able to log in. With login's limiter on the
// refresh route, the login below would get a 429.
func TestRefresh_DoesNotSpendLoginBudget(t *testing.T) {
hash, err := bcrypt.GenerateFromPassword([]byte("password"), bcryptCost)
require.NoError(t, err)
store := &loginStubStore{user: &User{
ID: 7,
Email: "driver@test.com",
PasswordHash: string(hash),
Role: "driver",
Active: true,
}}
loginLimiter := NewLoginRateLimiter()
defer loginLimiter.Stop()
refreshLimiter := NewRefreshRateLimiter()
defer refreshLimiter.Stop()
mux := newMux(store, nil, nil, testSecret, testTTLs, time.Time{}, loginLimiter, refreshLimiter, false, false, nil, nil, nil)

var w *httptest.ResponseRecorder
for range refreshIPLimit + 1 {
w = postRefresh(mux.ServeHTTP, "unknown-token")
}
require.Equal(t, http.StatusTooManyRequests, w.Code, "the refresh budget must be spent before login is tried")

login := postLogin(mux.ServeHTTP, "driver@test.com", "password")
assert.Equal(t, http.StatusOK, login.Code, "spending the refresh budget must not spend login's")
assert.NotEmpty(t, decodeTokens(t, login).AccessToken)
}

func TestHandleRefresh_MalformedBody(t *testing.T) {
f, _ := refreshStoreWithUser(t)
handler := handleRefreshToken(f, testSecret, testTTLs, nil, false)
Expand Down
2 changes: 1 addition & 1 deletion handler_composition_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ func newTestHandler(t *testing.T, enabled bool) http.Handler {
t.Cleanup(tracker.Stop)
ll := NewLoginRateLimiter()
t.Cleanup(ll.Stop)
h, err := newHandler(&noopStore{}, tracker, nil, ll, testSecret, testTTLs, time.Now(), adminUIConfig{enabled: enabled, stalenessThreshold: 5 * time.Minute}, false, nil, nil, nil)
h, err := newHandler(&noopStore{}, tracker, nil, ll, nil, testSecret, testTTLs, time.Now(), adminUIConfig{enabled: enabled, stalenessThreshold: 5 * time.Minute}, false, nil, nil, nil)
require.NoError(t, err)
return h
}
Expand Down
13 changes: 8 additions & 5 deletions main.go
Original file line number Diff line number Diff line change
Expand Up @@ -78,15 +78,15 @@ type appStore interface {
// Extracting route registration here allows tests to build the real mux
// without a live database, catching middleware wiring gaps like the one fixed
// in issue #82.
func newMux(store appStore, tracker *Tracker, rateLimiter *VehicleRateLimiter, jwtSecret []byte, ttls tokenTTLs, startTime time.Time, loginLimiter *LoginRateLimiter, trustProxy, feedAuthEnabled bool, riderSvc *riderService, catalog *gtfsCatalog, mapPMTilesURL *string) *http.ServeMux {
func newMux(store appStore, tracker *Tracker, rateLimiter *VehicleRateLimiter, jwtSecret []byte, ttls tokenTTLs, startTime time.Time, loginLimiter *LoginRateLimiter, refreshLimiter *RefreshRateLimiter, trustProxy, feedAuthEnabled bool, riderSvc *riderService, catalog *gtfsCatalog, mapPMTilesURL *string) *http.ServeMux {
mux := http.NewServeMux()

authMiddleware := requireAuth(jwtSecret, store)
adminMiddleware := requireAdmin()
riderEstimates, riderStatus := riderOrOff(riderSvc)

mux.Handle("POST /api/v1/auth/login", handleLogin(store, store, jwtSecret, ttls, loginLimiter, trustProxy))
mux.Handle("POST /api/v1/auth/refresh", handleRefreshToken(store, jwtSecret, ttls, loginLimiter, trustProxy))
mux.Handle("POST /api/v1/auth/refresh", handleRefreshToken(store, jwtSecret, ttls, refreshLimiter, trustProxy))
mux.Handle("POST /api/v1/auth/logout", authMiddleware(handleLogout(store, store)))
feed := handleGetFeed(tracker, riderEstimates)
if feedAuthEnabled {
Expand Down Expand Up @@ -159,10 +159,10 @@ func newMux(store appStore, tracker *Tracker, rateLimiter *VehicleRateLimiter, j
// CSRF protection wrapping the whole thing. It is the single place routes
// and cross-cutting middleware come together.
func newHandler(store appStore, tracker *Tracker, rateLimiter *VehicleRateLimiter,
loginLimiter *LoginRateLimiter, jwtSecret []byte, ttls tokenTTLs, startTime time.Time,
loginLimiter *LoginRateLimiter, refreshLimiter *RefreshRateLimiter, jwtSecret []byte, ttls tokenTTLs, startTime time.Time,
cfg adminUIConfig, feedAuthEnabled bool, riderSvc *riderService, catalog *gtfsCatalog, mapPMTilesURL *string) (http.Handler, error) {

mux := newMux(store, tracker, rateLimiter, jwtSecret, ttls, startTime, loginLimiter, cfg.trustProxy, feedAuthEnabled, riderSvc, catalog, mapPMTilesURL)
mux := newMux(store, tracker, rateLimiter, jwtSecret, ttls, startTime, loginLimiter, refreshLimiter, cfg.trustProxy, feedAuthEnabled, riderSvc, catalog, mapPMTilesURL)

if cfg.enabled {
ui, err := newAdminUI(store, tracker, jwtSecret, loginLimiter, cfg)
Expand Down Expand Up @@ -278,6 +278,9 @@ func main() {
loginLimiter := NewLoginRateLimiter()
defer loginLimiter.Stop()

refreshLimiter := NewRefreshRateLimiter()
defer refreshLimiter.Stop()

// Expired refresh tokens are unusable by every path that reads them, so
// pruning them is garbage collection rather than a retention policy — it
// runs by default and has no "keep forever" setting. A bad interval only
Expand Down Expand Up @@ -353,7 +356,7 @@ func main() {

startTime := time.Now()

handler, err := newHandler(store, tracker, rateLimiter, loginLimiter, jwtSecret, ttls, startTime,
handler, err := newHandler(store, tracker, rateLimiter, loginLimiter, refreshLimiter, jwtSecret, ttls, startTime,
adminUIConfig{enabled: adminUIEnabled(), trustProxy: trustProxyHeaders(), stalenessThreshold: maxAge}, feedAuthEnabled, riderSvc, catalog, mapPMTilesURL)
if err != nil {
slog.Error("failed to build handler", "error", err)
Expand Down
6 changes: 3 additions & 3 deletions map_handlers_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -138,7 +138,7 @@ func TestMapRoute_Wiring(t *testing.T) {
riderToken, err := generateRiderJWT("rider-1", testSecret, time.Hour)
require.NoError(t, err)
mapURL := testMapURL
mux := newMux(&noopStore{}, nil, nil, testSecret, testTTLs, time.Time{}, nil, false, false, nil, nil, &mapURL)
mux := newMux(&noopStore{}, nil, nil, testSecret, testTTLs, time.Time{}, nil, nil, false, false, nil, nil, &mapURL)

tests := []struct {
name string
Expand Down Expand Up @@ -173,7 +173,7 @@ func TestMapRoute_Wiring(t *testing.T) {
func TestMapRoute_RegisteredWithoutAMapFile(t *testing.T) {
driverToken, err := generateJWT(&User{ID: 1, Email: "driver@test.com", Role: "driver"}, testSecret, defaultAccessTokenTTL)
require.NoError(t, err)
mux := newMux(&noopStore{}, nil, nil, testSecret, testTTLs, time.Time{}, nil, false, false, nil, nil, nil)
mux := newMux(&noopStore{}, nil, nil, testSecret, testTTLs, time.Time{}, nil, nil, false, false, nil, nil, nil)

req := httptest.NewRequest(http.MethodGet, "/api/v1/map", nil)
req.Header.Set("Authorization", "Bearer "+driverToken)
Expand All @@ -194,7 +194,7 @@ func TestNewHandler_PassesTheMapURLOn(t *testing.T) {
driverToken, err := generateJWT(&User{ID: 1, Email: "driver@test.com", Role: "driver"}, testSecret, defaultAccessTokenTTL)
require.NoError(t, err)
mapURL := testMapURL
h, err := newHandler(&noopStore{}, tracker, nil, loginLimiter, testSecret, testTTLs, time.Now(),
h, err := newHandler(&noopStore{}, tracker, nil, loginLimiter, nil, testSecret, testTTLs, time.Now(),
adminUIConfig{stalenessThreshold: 5 * time.Minute}, false, nil, nil, &mapURL)
require.NoError(t, err)

Expand Down
6 changes: 3 additions & 3 deletions openapi.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -166,9 +166,9 @@ paths:
$ref: '#/components/responses/UnsupportedMediaType'
'429':
description: |
Too many attempts from this IP. The endpoint shares the login
limiter's per-IP budget, so attempts spread across both auth
endpoints draw on one allowance.
Too many attempts from this IP. The endpoint has its own per-IP
budget, separate from login's, so a burst of refreshes does not
block logins from the same address.
content:
application/json:
schema:
Expand Down
16 changes: 3 additions & 13 deletions ratelimit_login.go
Original file line number Diff line number Diff line change
Expand Up @@ -58,17 +58,6 @@ func (l *LoginRateLimiter) Allow(ip, email string) bool {
return allowInWindow(l.byEmail, email, loginEmailLimit, now, "login")
}

// AllowIP applies only the per-IP window. The refresh endpoint uses it
// because a caller presenting a refresh token has no email to key the second
// dimension on. Sharing this limiter's per-IP budget with login is
// deliberate: an address gets one attempt allowance across both auth
// endpoints, not one each.
func (l *LoginRateLimiter) AllowIP(ip string) bool {
l.mu.Lock()
defer l.mu.Unlock()
return allowInWindow(l.byIP, ip, loginIPLimit, time.Now(), "login")
}

// ResetEmail clears the per-email window after a successful authentication,
// so an account legitimately signing in several times a minute (a shared
// account, or the driver app plus the admin form) isn't 429'd despite zero
Expand All @@ -81,8 +70,9 @@ func (l *LoginRateLimiter) ResetEmail(email string) {
delete(l.byEmail, email)
}

// allowInWindow is the fixed-window admission shared by the login and rider
// registration limiters; name says which one is speaking when it fails closed.
// allowInWindow is the fixed-window admission shared by the login, refresh and
// rider registration limiters; name says which one is speaking when it fails
// closed.
func allowInWindow(m map[string]*loginWindowEntry, key string, limit int, now time.Time, name string) bool {
e, ok := m[key]
if !ok {
Expand Down
41 changes: 0 additions & 41 deletions ratelimit_login_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -84,44 +84,3 @@ func TestLoginRateLimiterEmailMapCapacityFailsClosed(t *testing.T) {

assert.False(t, l.Allow("9.9.9.9", "newcomer@test.com"), "new email denied when byEmail is at capacity")
}

func TestLoginRateLimiterAllowIP(t *testing.T) {
l := NewLoginRateLimiter()
defer l.Stop()

for i := range loginIPLimit {
assert.True(t, l.AllowIP("1.2.3.4"), "attempt %d", i)
}
assert.False(t, l.AllowIP("1.2.3.4"), "the per-IP window must close after loginIPLimit attempts")
assert.True(t, l.AllowIP("5.6.7.8"), "other IPs are unaffected")
}

// TestLoginRateLimiterAllowIPSharesBudgetWithLogin pins the decision to reuse
// this limiter for the refresh endpoint rather than adding a second one: an
// address gets one allowance across both auth endpoints, so it cannot double
// its attempts by alternating between them.
func TestLoginRateLimiterAllowIPSharesBudgetWithLogin(t *testing.T) {
l := NewLoginRateLimiter()
defer l.Stop()

for i := range loginIPLimit {
assert.True(t, l.AllowIP("1.2.3.4"), "refresh attempt %d", i)
}
assert.False(t, l.Allow("1.2.3.4", "driver@test.com"),
"refresh attempts must consume the same per-IP budget login uses")
}

// TestLoginRateLimiterAllowIPLeavesEmailBudget verifies AllowIP touches only
// the IP dimension: a refresh from one address must not spend an unrelated
// account's per-email allowance.
func TestLoginRateLimiterAllowIPLeavesEmailBudget(t *testing.T) {
l := NewLoginRateLimiter()
defer l.Stop()

for range loginIPLimit {
assert.True(t, l.AllowIP("1.2.3.4"))
}
for i := range loginEmailLimit {
assert.True(t, l.Allow("5.6.7.8", "driver@test.com"), "attempt %d", i)
}
}
67 changes: 67 additions & 0 deletions ratelimit_refresh.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
package main

import (
"sync"
"time"
)

// refreshIPLimit is the number of token refreshes allowed per IP per
// loginWindow. Drivers who log in at the same pull-out get access tokens that
// expire together, so they refresh together, and a depot's drivers usually
// share one public IP. 60 covers a 50-driver depot's burst with room left for
// retries; a bigger depot behind one address would still be throttled.
const refreshIPLimit = 60

// RefreshRateLimiter guards token refresh with a per-IP fixed window. Like
// LoginRateLimiter it FAILS CLOSED at capacity, since refresh mints
// credentials too.
//
// It does not share LoginRateLimiter's per-IP budget. Login's budget is tight
// because a password can be guessed; a refresh token is 32 random bytes and
// cannot be, so a separate allowance gives an attacker nothing against login.
// What this limiter guards against is flooding, and that needs a cap generous
// enough for a depot's synchronised refreshes not to lock its drivers out of
// logging in.
type RefreshRateLimiter struct {
mu sync.Mutex
byIP map[string]*loginWindowEntry
limit int
stop chan struct{}
once sync.Once
}

func NewRefreshRateLimiter() *RefreshRateLimiter {
l := &RefreshRateLimiter{
byIP: make(map[string]*loginWindowEntry),
limit: refreshIPLimit,
stop: make(chan struct{}),
}
go l.cleanup()
return l
}

func (l *RefreshRateLimiter) Stop() { l.once.Do(func() { close(l.stop) }) }

// Allow reports whether another refresh from ip fits inside the current
// window.
func (l *RefreshRateLimiter) Allow(ip string) bool {
l.mu.Lock()
defer l.mu.Unlock()
return allowInWindow(l.byIP, ip, l.limit, time.Now(), "refresh")
}

func (l *RefreshRateLimiter) cleanup() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for {
select {
case <-ticker.C:
cutoff := time.Now().Add(-2 * loginWindow)
l.mu.Lock()
pruneStaleWindows(l.byIP, cutoff)
l.mu.Unlock()
case <-l.stop:
return
}
}
}
Loading
Loading