diff --git a/cmd/strelaysrv/listener.go b/cmd/strelaysrv/listener.go index 164769047..3c4a2e427 100644 --- a/cmd/strelaysrv/listener.go +++ b/cmd/strelaysrv/listener.go @@ -184,7 +184,7 @@ func protocolConnectionHandler(tcpConn net.Conn, config *tls.Config, token strin continue } // requestedPeer is the server, id is the client - ses := newSession(requestedPeer, id, sessionLimiter, globalLimiter) + ses := newSession(requestedPeer, id, sessionLimitBps, globalLimiter) go ses.Serve() diff --git a/cmd/strelaysrv/main.go b/cmd/strelaysrv/main.go index e96e9af86..c1d264081 100644 --- a/cmd/strelaysrv/main.go +++ b/cmd/strelaysrv/main.go @@ -51,7 +51,6 @@ var ( globalLimitBps int overLimit atomic.Bool descriptorLimit int64 - sessionLimiter *rate.Limiter globalLimiter *rate.Limiter networkBufferSize int @@ -228,9 +227,6 @@ func main() { } } - if sessionLimitBps > 0 { - sessionLimiter = rate.NewLimiter(rate.Limit(sessionLimitBps), 2*sessionLimitBps) - } if globalLimitBps > 0 { globalLimiter = rate.NewLimiter(rate.Limit(globalLimitBps), 2*globalLimitBps) } diff --git a/cmd/strelaysrv/session.go b/cmd/strelaysrv/session.go index 1426216e5..79d1184fc 100644 --- a/cmd/strelaysrv/session.go +++ b/cmd/strelaysrv/session.go @@ -27,7 +27,7 @@ var ( bytesProxied atomic.Int64 ) -func newSession(serverid, clientid syncthingprotocol.DeviceID, sessionRateLimit, globalRateLimit *rate.Limiter) *session { +func newSession(serverid, clientid syncthingprotocol.DeviceID, sessionLimitBps int, globalRateLimit *rate.Limiter) *session { serverkey := make([]byte, 32) _, err := rand.Read(serverkey) if err != nil { @@ -40,12 +40,17 @@ func newSession(serverid, clientid syncthingprotocol.DeviceID, sessionRateLimit, return nil } + var sessionRateLimit *rate.Limiter + if sessionLimitBps > 0 { + sessionRateLimit = rate.NewLimiter(rate.Limit(sessionLimitBps), 2*sessionLimitBps) + } ses := &session{ serverkey: serverkey, serverid: serverid, clientkey: clientkey, clientid: clientid, rateLimit: makeRateLimitFunc(sessionRateLimit, globalRateLimit), + limiter: sessionRateLimit, connsChan: make(chan net.Conn), conns: make([]net.Conn, 0, 2), } @@ -109,6 +114,7 @@ type session struct { clientid syncthingprotocol.DeviceID rateLimit func(bytes int) + limiter *rate.Limiter connsChan chan net.Conn conns []net.Conn