diff --git a/server/middlewares.go b/server/middlewares.go index c54e947f3..674337e92 100644 --- a/server/middlewares.go +++ b/server/middlewares.go @@ -232,15 +232,20 @@ func trustedProxyPrefixes(list string) []string { return prefixes } -// ClientIPRateLimiter returns a rate limiter keyed by the client IP resolved by realIPMiddleware, -// so spoofed forwarding headers cannot be rotated for a fresh bucket. It falls back to the peer -// address, so that a missing middleware degrades to per-peer limiting rather than one shared bucket. +// ClientIPRateLimiter returns a rate limiter keyed by ClientIP, so spoofed forwarding headers +// cannot be rotated for a fresh bucket. func ClientIPRateLimiter(requestLimit int, windowLength time.Duration) func(http.Handler) http.Handler { return httprate.LimitBy(requestLimit, windowLength, func(r *http.Request) (string, error) { - return httprate.CanonicalizeIP(cmp.Or(middleware.GetClientIP(r.Context()), peerHost(r))), nil + return ClientIP(r), nil }) } +// ClientIP returns the canonical client IP resolved by realIPMiddleware, for keying rate limits. The +// peer address fallback degrades a missing middleware to per-peer limiting, not one shared bucket. +func ClientIP(r *http.Request) string { + return httprate.CanonicalizeIP(cmp.Or(middleware.GetClientIP(r.Context()), peerHost(r))) +} + // reqToCtx creates a middleware that updates the request's context with a value computed from the request. A given key // can only be set once. func reqToCtx(key any, fn func(req *http.Request) any) func(http.Handler) http.Handler { diff --git a/server/subsonic/auth_limiter.go b/server/subsonic/auth_limiter.go new file mode 100644 index 000000000..88df6eb50 --- /dev/null +++ b/server/subsonic/auth_limiter.go @@ -0,0 +1,134 @@ +package subsonic + +import ( + "cmp" + "context" + "hash/maphash" + "sync" + "time" + + "github.com/navidrome/navidrome/consts" +) + +// authLimiter caps failed Subsonic logins per key. Checks run at most `limit` at a time and failures +// are recorded afterwards, so a window admits up to 2*limit-1 guesses and valid requests only wait. +type authLimiter struct { + limit int + window time.Duration + seed maphash.Seed + mu sync.Mutex + keys map[uint64]*authAttempts // hashed, so attacker-chosen usernames cannot bloat memory + lastSweep time.Time +} + +type authAttempts struct { + failures int + start time.Time + slots chan struct{} + refs int +} + +// authSlot is a reserved credential check. A nil slot releases nothing, which is what a disabled +// limiter hands back. +type authSlot struct { + limiter *authLimiter + entry *authAttempts +} + +// newAuthLimiter returns nil when limit is not positive. A nil limiter allows everything. +func newAuthLimiter(limit int, window time.Duration) *authLimiter { + if limit <= 0 { + return nil + } + return &authLimiter{ + limit: limit, + window: cmp.Or(window, consts.DefaultAuthWindowLength), + seed: maphash.MakeSeed(), + keys: map[uint64]*authAttempts{}, + } +} + +// acquire reserves a credential check for key, waiting while other checks for the same key are in +// flight. It only fails when the key already reached `limit` failures in the current window. +func (l *authLimiter) acquire(ctx context.Context, key string) (*authSlot, bool) { + if l == nil { + return nil, true + } + a, ok := l.reserve(key) + if !ok { + return nil, false + } + + select { + case a.slots <- struct{}{}: + case <-ctx.Done(): + l.unref(a) + return nil, false + } + + l.mu.Lock() + blocked := a.failures >= l.limit + if blocked { + a.refs-- + } + l.mu.Unlock() + if blocked { + <-a.slots + return nil, false + } + return &authSlot{limiter: l, entry: a}, true +} + +func (l *authLimiter) reserve(key string) (*authAttempts, bool) { + now := time.Now() + l.mu.Lock() + defer l.mu.Unlock() + l.sweep(now) + + h := maphash.String(l.seed, key) + a := l.keys[h] + switch { + case a == nil: + a = &authAttempts{start: now, slots: make(chan struct{}, l.limit)} + l.keys[h] = a + case now.Sub(a.start) >= l.window: + a.failures, a.start = 0, now + } + if a.failures >= l.limit { + return nil, false + } + a.refs++ + return a, true +} + +func (l *authLimiter) unref(a *authAttempts) { + l.mu.Lock() + a.refs-- + l.mu.Unlock() +} + +func (s *authSlot) release(failed bool) { + if s == nil { + return + } + s.limiter.mu.Lock() + if failed { + s.entry.failures++ + } + s.entry.refs-- + s.limiter.mu.Unlock() + <-s.entry.slots +} + +// sweep drops idle expired keys once per window, so memory is bounded by recent attempts. +func (l *authLimiter) sweep(now time.Time) { + if now.Sub(l.lastSweep) < l.window { + return + } + for h, a := range l.keys { + if a.refs == 0 && now.Sub(a.start) >= l.window { + delete(l.keys, h) + } + } + l.lastSweep = now +} diff --git a/server/subsonic/auth_limiter_test.go b/server/subsonic/auth_limiter_test.go new file mode 100644 index 000000000..66e1a594f --- /dev/null +++ b/server/subsonic/auth_limiter_test.go @@ -0,0 +1,143 @@ +package subsonic + +import ( + "context" + "sync" + "sync/atomic" + "testing" + "testing/synctest" + "time" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("authLimiter", func() { + var ctx context.Context + + BeforeEach(func() { + ctx = context.Background() + }) + + acquire := func(l *authLimiter, key string) (*authSlot, bool) { + GinkgoHelper() + return l.acquire(ctx, key) + } + + It("blocks a key after the configured number of failures", func() { + l := newAuthLimiter(2, time.Minute) + for range 2 { + slot, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + slot.release(true) + } + + _, ok := acquire(l, "k") + Expect(ok).To(BeFalse()) + }) + + It("never counts successful checks", func() { + l := newAuthLimiter(2, time.Minute) + for range 50 { + slot, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + slot.release(false) + } + }) + + It("keeps keys independent", func() { + l := newAuthLimiter(1, time.Minute) + slot, _ := acquire(l, "a") + slot.release(true) + _, ok := acquire(l, "a") + Expect(ok).To(BeFalse()) + + _, ok = acquire(l, "b") + Expect(ok).To(BeTrue()) + }) + + It("waits for an in-flight check instead of failing the request", func() { + l := newAuthLimiter(1, time.Minute) + held, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + + waiting := make(chan bool, 1) + go func() { + slot, ok := l.acquire(ctx, "k") + slot.release(false) + waiting <- ok + }() + Consistently(waiting, 50*time.Millisecond).ShouldNot(Receive()) + + held.release(false) + Eventually(waiting).Should(Receive(BeTrue())) + }) + + It("stops waiting when the request is canceled", func() { + l := newAuthLimiter(1, time.Minute) + held, _ := acquire(l, "k") + DeferCleanup(func() { held.release(false) }) + + canceled, cancel := context.WithCancel(context.Background()) + cancel() + _, ok := l.acquire(canceled, "k") + Expect(ok).To(BeFalse()) + }) + + It("does not let concurrent guesses overshoot the limit", func() { + l := newAuthLimiter(5, time.Minute) + hold := make(chan struct{}) + var checks atomic.Int32 + var wg sync.WaitGroup + for range 50 { + wg.Go(func() { + slot, ok := l.acquire(ctx, "k") + if !ok { + return + } + checks.Add(1) + <-hold + slot.release(true) + }) + } + + Eventually(checks.Load).Should(Equal(int32(5))) + Consistently(checks.Load, 100*time.Millisecond).Should(Equal(int32(5))) + close(hold) + wg.Wait() + + _, ok := acquire(l, "k") + Expect(ok).To(BeFalse()) + }) + + It("allows everything when the limit is disabled", func() { + l := newAuthLimiter(0, time.Minute) + for range 10 { + slot, ok := acquire(l, "k") + Expect(ok).To(BeTrue()) + slot.release(true) + } + }) +}) + +// testing/synctest's fake clock needs a *testing.T, which Ginkgo doesn't give. +func TestAuthLimiterWindow(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + g := NewWithT(t) + ctx := context.Background() + l := newAuthLimiter(1, 20*time.Second) + for _, key := range []string{"a", "b"} { + slot, ok := l.acquire(ctx, key) + g.Expect(ok).To(BeTrue()) + slot.release(true) + } + _, ok := l.acquire(ctx, "a") + g.Expect(ok).To(BeFalse()) + + time.Sleep(20 * time.Second) + slot, ok := l.acquire(ctx, "a") + g.Expect(ok).To(BeTrue()) + slot.release(false) + g.Expect(l.keys).To(HaveLen(1), "expired keys must be dropped") + }) +} diff --git a/server/subsonic/middlewares.go b/server/subsonic/middlewares.go index 6dfa2263f..35e13eaa5 100644 --- a/server/subsonic/middlewares.go +++ b/server/subsonic/middlewares.go @@ -98,6 +98,7 @@ func checkRequiredParameters(next http.Handler) http.Handler { } func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { + limiter := newAuthLimiter(conf.Server.AuthRequestLimit, conf.Server.AuthWindowLength) return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { ctx := r.Context() @@ -126,21 +127,32 @@ func authenticate(ds model.DataStore) func(next http.Handler) http.Handler { salt, _ := p.String("s") jwt, _ := p.String("jwt") - usr, err = ds.User(ctx).FindByUsernameWithPassword(username) - if errors.Is(err, context.Canceled) { - log.Debug(ctx, "API: Request canceled when authenticating", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) + // Blocked attempts get the same response as a wrong password, so they reveal nothing + limitKey := server.ClientIP(r) + "\x00" + strings.ToLower(username) + slot, allowed := limiter.acquire(ctx, limitKey) + if !allowed { + if ctx.Err() != nil { + return + } + log.Warn(ctx, "API: Too many failed login attempts", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr) + sendError(w, r, newError(responses.ErrorAuthenticationFail)) return } + + usr, err = ds.User(ctx).FindByUsernameWithPassword(username) + if err == nil { + err = validateCredentials(usr, pass, token, salt, jwt) + } + invalidLogin := errors.Is(err, model.ErrNotFound) || errors.Is(err, model.ErrInvalidAuth) + slot.release(invalidLogin) switch { - case errors.Is(err, model.ErrNotFound): + case errors.Is(err, context.Canceled): + log.Debug(ctx, "API: Request canceled when authenticating", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) + return + case invalidLogin: log.Warn(ctx, "API: Invalid login", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) case err != nil: log.Error(ctx, "API: Error authenticating username", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) - default: - err = validateCredentials(usr, pass, token, salt, jwt) - if err != nil { - log.Warn(ctx, "API: Invalid login", "auth", "subsonic", "username", username, "remoteAddr", r.RemoteAddr, err) - } } } diff --git a/server/subsonic/middlewares_test.go b/server/subsonic/middlewares_test.go index cb34b92e7..62de09e76 100644 --- a/server/subsonic/middlewares_test.go +++ b/server/subsonic/middlewares_test.go @@ -8,6 +8,8 @@ import ( "net/http" "net/http/httptest" "strings" + "sync" + "sync/atomic" "time" "github.com/navidrome/navidrome/conf" @@ -306,6 +308,166 @@ var _ = Describe("Middlewares", func() { Expect(next.called).To(BeFalse()) }) }) + + When("failed attempts reach AuthRequestLimit", func() { + var cp http.Handler + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.AuthRequestLimit = 3 + conf.Server.AuthWindowLength = time.Minute + cp = authenticate(ds)(next) + }) + + serve := func(r *http.Request) *httptest.ResponseRecorder { + next.called = false + rec := httptest.NewRecorder() + cp.ServeHTTP(rec, r) + return rec + } + failTimes := func(n int, params ...string) { + for range n { + Expect(serve(newGetRequest(params...)).Body.String()).To(ContainSubstring(`code="40"`)) + } + } + + It("rejects the correct password exactly like a wrong one", func() { + failTimes(3, "u=admin", "p=WRONG") + + rec := serve(newGetRequest("u=admin", "p=wordpass")) + + Expect(next.called).To(BeFalse()) + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(rec.Body.String()).To(ContainSubstring(`code="40"`)) + Expect(rec.Header().Get("Retry-After")).To(BeEmpty()) + }) + + It("counts attempts against unknown usernames", func() { + failTimes(3, "u=newuser", "p=secret") + _ = ds.User(context.TODO()).Put(&model.User{UserName: "newuser", NewPassword: "secret"}) + + serve(newGetRequest("u=newuser", "p=secret")) + Expect(next.called).To(BeFalse()) + }) + + It("treats usernames case-insensitively", func() { + failTimes(3, "u=ADMIN", "p=WRONG") + + serve(newGetRequest("u=admin", "p=wordpass")) + Expect(next.called).To(BeFalse()) + }) + + It("does not count successful logins", func() { + for range 10 { + serve(newGetRequest("u=admin", "p=wordpass")) + Expect(next.called).To(BeTrue()) + } + }) + + It("does not count server errors", func() { + userRepo := ds.User(context.TODO()).(*tests.MockedUserRepo) + userRepo.Error = errors.New("db down") + failTimes(5, "u=admin", "p=wordpass") + userRepo.Error = nil + + serve(newGetRequest("u=admin", "p=wordpass")) + Expect(next.called).To(BeTrue()) + }) + + It("does not block other usernames from the same IP", func() { + _ = ds.User(context.TODO()).Put(&model.User{UserName: "other", NewPassword: "otherpass"}) + failTimes(3, "u=admin", "p=WRONG") + + serve(newGetRequest("u=other", "p=otherpass")) + Expect(next.called).To(BeTrue()) + }) + + It("does not block the same username from another IP", func() { + failTimes(3, "u=admin", "p=WRONG") + + r := newGetRequest("u=admin", "p=wordpass") + r.RemoteAddr = "198.51.100.7:1234" + serve(r) + Expect(next.called).To(BeTrue()) + }) + + It("does not limit reverse proxy authentication", func() { + conf.Server.ExtAuth.TrustedSources = "192.168.1.1/24" + conf.Server.ExtAuth.UserHeader = "Remote-User" + failTimes(3, "u=admin", "p=WRONG") + + r := newGetRequest() + r.Header.Add("Remote-User", "admin") + r = r.WithContext(request.WithReverseProxyIp(r.Context(), "192.168.1.1")) + serve(r) + Expect(next.called).To(BeTrue()) + }) + + It("is disabled when AuthRequestLimit is 0", func() { + conf.Server.AuthRequestLimit = 0 + cp = authenticate(ds)(next) + failTimes(10, "u=admin", "p=WRONG") + + serve(newGetRequest("u=admin", "p=wordpass")) + Expect(next.called).To(BeTrue()) + }) + }) + + When("valid requests overlap", func() { + var gate *gatedUserRepo + var gatedDS model.DataStore + + BeforeEach(func() { + DeferCleanup(configtest.SetupConfig()) + conf.Server.AuthRequestLimit = 5 + conf.Server.AuthWindowLength = time.Minute + gate = &gatedUserRepo{ + UserRepository: ds.User(context.TODO()), + entered: make(chan struct{}, 64), + proceed: make(chan struct{}), + } + gatedDS = &gatedDataStore{DataStore: ds, users: gate} + }) + + It("lets every valid request through while checks are in flight", func() { + const burst = 6 + cp := authenticate(gatedDS)(&countingHandler{}) + var passed atomic.Int32 + var wg sync.WaitGroup + for range burst { + wg.Go(func() { + rec := httptest.NewRecorder() + cp.ServeHTTP(rec, newGetRequest("u=admin", "p=wordpass")) + if !strings.Contains(rec.Body.String(), `code="40"`) { + passed.Add(1) + } + }) + } + for range conf.Server.AuthRequestLimit { + Eventually(gate.entered).Should(Receive()) + } + close(gate.proceed) + wg.Wait() + + Expect(passed.Load()).To(Equal(int32(burst))) + }) + + It("caps concurrent credential checks for wrong passwords", func() { + cp := authenticate(gatedDS)(&countingHandler{}) + var wg sync.WaitGroup + for i := range 100 { + wg.Go(func() { + cp.ServeHTTP(httptest.NewRecorder(), newGetRequest("u=admin", fmt.Sprintf("p=wrong%d", i))) + }) + } + + limit := int32(conf.Server.AuthRequestLimit) + Eventually(gate.lookups.Load).Should(Equal(limit)) + Consistently(gate.lookups.Load, 100*time.Millisecond).Should(Equal(limit)) + close(gate.proceed) + wg.Wait() + }) + }) }) Describe("AdminOnly", func() { @@ -560,3 +722,28 @@ func (mp *mockPlayers) Register(ctx context.Context, id, client, typ, ip string) } return &model.Player{ID: id}, mp.transcoding, nil } + +type gatedDataStore struct { + model.DataStore + users model.UserRepository +} + +func (g *gatedDataStore) User(context.Context) model.UserRepository { return g.users } + +type gatedUserRepo struct { + model.UserRepository + entered chan struct{} + proceed chan struct{} + lookups atomic.Int32 +} + +func (g *gatedUserRepo) FindByUsernameWithPassword(username string) (*model.User, error) { + g.lookups.Add(1) + g.entered <- struct{}{} + <-g.proceed + return g.UserRepository.FindByUsernameWithPassword(username) +} + +type countingHandler struct{ calls atomic.Int32 } + +func (c *countingHandler) ServeHTTP(http.ResponseWriter, *http.Request) { c.calls.Add(1) }