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
57 changes: 4 additions & 53 deletions pkg/api/middleware/ratelimit.go
Original file line number Diff line number Diff line change
Expand Up @@ -26,18 +26,17 @@ import (
"time"

"github.com/ZaparooProject/zaparoo-core/v2/pkg/helpers/syncutil"
"github.com/olahol/melody"
"github.com/rs/zerolog/log"
"golang.org/x/time/rate"
)

const (
RequestsPerMinute = 100 // Simple limit - 100 requests per minute per IP
BurstSize = 20 // Allow burst of 20 requests
WebSocketRateLimitWait = 2 * time.Second
RequestsPerMinute = 100 // Simple limit - 100 requests per minute per IP
BurstSize = 20 // Allow burst of 20 requests
)

// IPRateLimiter manages rate limiters per IP address for both HTTP and WebSocket
// IPRateLimiter manages HTTP admission rate limiters per IP address.
// WebSocket upgrades consume one token; established frames use session queues.
type IPRateLimiter struct {
limiters map[string]*rateLimiterEntry
mu syncutil.RWMutex
Expand Down Expand Up @@ -174,51 +173,3 @@ func HTTPRateLimitMiddleware(limiter *IPRateLimiter) func(http.Handler) http.Han
})
}
}

// WebSocketRateLimitHandler wraps a WebSocket message handler with rate
// limiting. When the per-IP rate limit is exceeded the connection is closed
// rather than receiving a structured JSON-RPC error: this avoids leaking
// plaintext frames onto encrypted sessions (which would not match the
// {"e":...} envelope and could not be decrypted by the client) and gives
// well-behaved clients an unambiguous "back off and reconnect" signal.
func WebSocketRateLimitHandler(
limiter *IPRateLimiter,
handler func(*melody.Session, []byte),
) func(*melody.Session, []byte) {
return WebSocketRateLimitHandlerWithWait(limiter, WebSocketRateLimitWait, handler)
}

func WebSocketRateLimitHandlerWithWait(
limiter *IPRateLimiter,
waitTimeout time.Duration,
handler func(*melody.Session, []byte),
) func(*melody.Session, []byte) {
return func(session *melody.Session, msg []byte) {
host, exempt := remoteRateLimitHost(session.Request.RemoteAddr)
if exempt {
handler(session, msg)
return
}

rl := limiter.GetLimiter(host)

ctx, cancel := context.WithTimeout(context.Background(), waitTimeout)
defer cancel()
waitTimeoutValue := waitTimeout
if err := rl.Wait(ctx); err != nil {
log.Warn().
Err(err).
Str("ip", host).
Int("msg_size", len(msg)).
Str("wait_timeout", waitTimeoutValue.String()).
Msg("WebSocket rate limit wait failed, closing connection")

if err := session.Close(); err != nil {
log.Debug().Err(err).Msg("failed to close rate-limited session")
}
return
}

handler(session, msg)
}
}
141 changes: 0 additions & 141 deletions pkg/api/middleware/ratelimit_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,19 +21,12 @@ package middleware

import (
"context"
"errors"
"net"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"

"github.com/gorilla/websocket"
"github.com/olahol/melody"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.org/x/time/rate"
)

Expand Down Expand Up @@ -245,137 +238,3 @@ func TestHTTPRateLimitMiddleware_DoesNotExemptNonLoopbackHostnames(t *testing.T)
}
assert.Equal(t, 1, callCount)
}

func TestWebSocketRateLimitHandler_WaitsForToken(t *testing.T) {
t.Parallel()

// burst=1 allows the first message immediately; the second waits for
// the next token instead of closing the session.
rl := NewIPRateLimiterWithLimits(rate.Every(50*time.Millisecond), 1)
var handlerCalls atomic.Int32
inner := func(_ *melody.Session, _ []byte) {
handlerCalls.Add(1)
}
wrapped := WebSocketRateLimitHandlerWithWait(rl, 250*time.Millisecond, inner)

m := melody.New()
m.HandleMessage(wrapped)

srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
r.RemoteAddr = "192.168.1.1:12345"
_ = m.HandleRequest(w, r)
}))
defer srv.Close()

wsURL := "ws" + strings.TrimPrefix(srv.URL, "http")
//nolint:bodyclose // websocket conn manages the body
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
require.NoError(t, err)
defer func() { _ = conn.Close() }()

// First message should go through.
require.NoError(t, conn.WriteMessage(websocket.TextMessage, []byte("hello")))
require.Eventually(t, func() bool {
return handlerCalls.Load() == 1
}, 500*time.Millisecond, 10*time.Millisecond, "first message should be handled")

// Second message should wait for a token and then reach the handler.
require.NoError(t, conn.WriteMessage(websocket.TextMessage, []byte("again")))
require.Eventually(t, func() bool {
return handlerCalls.Load() == 2
}, 500*time.Millisecond, 10*time.Millisecond, "second message should be handled after backpressure wait")

// The server should keep the connection open.
require.NoError(t, conn.WriteMessage(websocket.TextMessage, []byte("third")))
require.Eventually(t, func() bool {
return handlerCalls.Load() == 3
}, 500*time.Millisecond, 10*time.Millisecond, "connection should remain usable after waiting")
}

func TestWebSocketRateLimitHandler_ClosesAfterWaitTimeout(t *testing.T) {
t.Parallel()

// burst=1 and no refill forces the second message to exceed the bounded wait.
rl := NewIPRateLimiterWithLimits(0, 1)
var handlerCalls atomic.Int32
inner := func(_ *melody.Session, _ []byte) {
handlerCalls.Add(1)
}
wrapped := WebSocketRateLimitHandlerWithWait(rl, 20*time.Millisecond, inner)

m := melody.New()
m.HandleMessage(wrapped)

srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
r.RemoteAddr = "192.168.1.1:12345"
_ = m.HandleRequest(w, r)
}))
defer srv.Close()

wsURL := "ws" + strings.TrimPrefix(srv.URL, "http")
//nolint:bodyclose // websocket conn manages the body
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
require.NoError(t, err)
defer func() { _ = conn.Close() }()

require.NoError(t, conn.WriteMessage(websocket.TextMessage, []byte("hello")))
require.Eventually(t, func() bool {
return handlerCalls.Load() == 1
}, 500*time.Millisecond, 10*time.Millisecond, "first message should be handled")

require.NoError(t, conn.WriteMessage(websocket.TextMessage, []byte("again")))
assert.Never(t, func() bool {
return handlerCalls.Load() != 1
}, 150*time.Millisecond, 10*time.Millisecond, "second message should not reach handler")

// The server should have closed the connection.
_ = conn.SetReadDeadline(time.Now().Add(100 * time.Millisecond))
_, _, err = conn.ReadMessage()
require.Error(t, err, "connection should be closed after rate limit exceeded")
var netErr net.Error
if errors.As(err, &netErr) {
assert.False(t, netErr.Timeout(), "connection read should fail from server close, not read timeout")
}
assert.True(t,
websocket.IsCloseError(err,
websocket.CloseNormalClosure,
websocket.CloseGoingAway,
websocket.CloseAbnormalClosure,
websocket.CloseNoStatusReceived,
) || websocket.IsUnexpectedCloseError(err),
"connection read should return a websocket close error, got %v", err,
)
}

func TestWebSocketRateLimitHandler_ExemptsLoopback(t *testing.T) {
t.Parallel()

rl := NewIPRateLimiterWithLimits(0, 1)
var handlerCalls atomic.Int32
inner := func(_ *melody.Session, _ []byte) {
handlerCalls.Add(1)
}
wrapped := WebSocketRateLimitHandlerWithWait(rl, 20*time.Millisecond, inner)

m := melody.New()
m.HandleMessage(wrapped)

srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_ = m.HandleRequest(w, r)
}))
defer srv.Close()

wsURL := "ws" + strings.TrimPrefix(srv.URL, "http")
//nolint:bodyclose // websocket conn manages the body
conn, _, err := websocket.DefaultDialer.Dial(wsURL, nil)
require.NoError(t, err)
defer func() { _ = conn.Close() }()

for range 3 {
require.NoError(t, conn.WriteMessage(websocket.TextMessage, []byte("hello")))
}

require.Eventually(t, func() bool {
return handlerCalls.Load() == 3
}, 500*time.Millisecond, 10*time.Millisecond, "loopback websocket messages should bypass rate limiting")
}
11 changes: 11 additions & 0 deletions pkg/api/request_priority.go
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,17 @@ const (
apiPriorityLow
)

func (p apiRequestPriority) String() string {
switch p {
case apiPriorityHigh:
return "high"
case apiPriorityLow:
return "low"
default:
return "normal"
}
}

func requestTimeoutForAPIMethod(method string) time.Duration {
if models.MethodHasUnboundedRuntime(method) {
return 0
Expand Down
43 changes: 36 additions & 7 deletions pkg/api/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -103,6 +103,11 @@ var JSONRPCErrorInternalError = models.ErrorObject{
Message: "Internal error",
}

var JSONRPCErrorServerBusy = models.ErrorObject{
Code: -32000,
Message: "Server busy",
}

func makeJSONRPCError(code int, message string) models.ErrorObject {
return models.ErrorObject{
Code: code,
Expand Down Expand Up @@ -1194,6 +1199,33 @@ func handleWSMessage(
}

if err := enqueueWSRequest(dispatcher, methodMap, &env, plaintext, cs, tracker); err != nil {
var queueFullErr *wsRequestQueueFullError
if errors.As(err, &queueFullErr) {
log.Warn().
Str("method", queueFullErr.method).
Str("requestId", requestIDForLog(queueFullErr.requestID)).
Str("priority", queueFullErr.priority.String()).
Int("queueDepth", queueFullErr.depth).
Int("queueCapacity", queueFullErr.capacity).
Msg("websocket request rejected because queue is full")
if queueFullErr.requestID.IsAbsent() {
endTrackedRequest()
return
}
dispatcher.enqueueResponse(&wsResponseJob{
result: requestResult{
ID: queueFullErr.requestID,
Error: &JSONRPCErrorServerBusy,
ShouldReply: true,
},
cs: cs,
tracker: tracker,
method: queueFullErr.method,
})
handoffTrackedRequest()
return
}

log.Warn().Err(err).Msg("failed to queue websocket request")
endTrackedRequest()
if sendErr := sendWSEncryptedError(
Expand Down Expand Up @@ -1883,13 +1915,10 @@ func StartWithReady(
r.Get("/api/v0.1/events", sseHandler)
})

session.HandleMessage(apimiddleware.WebSocketRateLimitHandler(
rateLimiter,
handleWSMessage(
methodMap, platform, cfg, st, inTokenQueue, confirmQueue,
db, limitsManager, profilesSvc, player, playbackManager, indexPauser, scrapePauser, backupPauser,
encGateway, lastSeenTracker, tracker,
),
session.HandleMessage(handleWSMessage(
methodMap, platform, cfg, st, inTokenQueue, confirmQueue,
db, limitsManager, profilesSvc, player, playbackManager, indexPauser, scrapePauser, backupPauser,
encGateway, lastSeenTracker, tracker,
))

// Static app assets
Expand Down
Loading
Loading