aboutsummaryrefslogtreecommitdiff
path: root/hotline/server.go
diff options
context:
space:
mode:
authorJeff Halter <868228+jhalter@users.noreply.github.com>2024-07-23 17:52:10 -0700
committerJeff Halter <868228+jhalter@users.noreply.github.com>2024-07-23 17:54:36 -0700
commit23c1295a4d9b13fdb955d2041a49b9b5d6daa060 (patch)
tree13baf607d05e1fde8a150b304342d44b79e1a7d2 /hotline/server.go
parent318eb6bc17ee406e26c1cfe0835409fd446bad75 (diff)
Add client connection rate limit
Diffstat (limited to 'hotline/server.go')
-rw-r--r--hotline/server.go26
1 files changed, 25 insertions, 1 deletions
diff --git a/hotline/server.go b/hotline/server.go
index 5531bd3..758f7ff 100644
--- a/hotline/server.go
+++ b/hotline/server.go
@@ -9,6 +9,7 @@ import (
"errors"
"fmt"
"golang.org/x/text/encoding/charmap"
+ "golang.org/x/time/rate"
"io"
"log"
"log/slog"
@@ -37,6 +38,8 @@ type Server struct {
NetInterface string
Port int
+ rateLimiters map[string]*rate.Limiter
+
handlers map[TranType]HandlerFunc
Config Config
@@ -98,6 +101,7 @@ func NewServer(options ...Option) (*Server, error) {
server := Server{
handlers: make(map[TranType]HandlerFunc),
outbox: make(chan Transaction),
+ rateLimiters: make(map[string]*rate.Limiter),
FS: &OSFileStore{},
ChatMgr: NewMemChatManager(),
ClientMgr: NewMemClientMgr(),
@@ -202,6 +206,10 @@ func (s *Server) processOutbox() {
}
}
+// perIPRateLimit controls how frequently an IP address can connect before being throttled.
+// 0.5 = 1 connection every 2 seconds
+const perIPRateLimit = rate.Limit(0.5)
+
func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
for {
select {
@@ -216,13 +224,29 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
}
go func() {
+ ipAddr := strings.Split(conn.RemoteAddr().(*net.TCPAddr).String(), ":")[0]
+
connCtx := context.WithValue(ctx, contextKeyReq, requestCtx{
remoteAddr: conn.RemoteAddr().String(),
})
- s.Logger.Info("Connection established", "addr", conn.RemoteAddr())
+ s.Logger.Info("Connection established", "ip", ipAddr)
defer conn.Close()
+ // Check if we have an existing rate limit for the IP and create one if we do not.
+ rl, ok := s.rateLimiters[ipAddr]
+ if !ok {
+ rl = rate.NewLimiter(perIPRateLimit, 1)
+ s.rateLimiters[ipAddr] = rl
+ }
+
+ // Check if the rate limit is exceeded and close the connection if so.
+ if !rl.Allow() {
+ s.Logger.Info("Rate limit exceeded", "RemoteAddr", conn.RemoteAddr())
+ conn.Close()
+ return
+ }
+
if err := s.handleNewConnection(connCtx, conn, conn.RemoteAddr().String()); err != nil {
if err == io.EOF {
s.Logger.Info("Client disconnected", "RemoteAddr", conn.RemoteAddr())