diff --git a/protocol/limiter/traffic/conn.go b/protocol/limiter/traffic/conn.go index 97e825b5..26c4bb63 100644 --- a/protocol/limiter/traffic/conn.go +++ b/protocol/limiter/traffic/conn.go @@ -20,14 +20,19 @@ func newConnWithUploadTrafficLimiter(ctx context.Context, conn net.Conn, limiter } func (conn *connWithTrafficLimiter) Write(p []byte) (int, error) { - reserved, err := conn.limiter.Reserve(uint64(len(p))) - if reserved < uint64(len(p)) { - conn.limiter.Commit(reserved, 0) + err := conn.limiter.Can(uint64(len(p))) + if err != nil { return 0, err } n, err := conn.Conn.Write(p) - conn.limiter.Commit(reserved, uint64(n)) - return n, err + if err != nil { + return 0, err + } + err = conn.limiter.Add(uint64(n)) + if err != nil { + return 0, err + } + return n, nil } type connWithUploadTrafficLimiter struct { @@ -37,16 +42,19 @@ type connWithUploadTrafficLimiter struct { } func (conn *connWithUploadTrafficLimiter) Read(p []byte) (int, error) { - reserved, err := conn.limiter.Reserve(uint64(len(p))) - if reserved == 0 { + err := conn.limiter.Can(1) + if err != nil { return 0, err } - if reserved < uint64(len(p)) { - p = p[:reserved] - } n, err := conn.Conn.Read(p) - conn.limiter.Commit(reserved, uint64(n)) - return n, err + if err != nil { + return 0, err + } + err = conn.limiter.Add(uint64(n)) + if err != nil { + return 0, err + } + return n, nil } type packetConnWithTrafficLimiter struct { @@ -64,14 +72,19 @@ func newPacketConnWithUploadTrafficLimiter(ctx context.Context, conn net.PacketC } func (conn *packetConnWithTrafficLimiter) WriteTo(p []byte, addr net.Addr) (int, error) { - reserved, err := conn.limiter.Reserve(uint64(len(p))) - if reserved < uint64(len(p)) { - conn.limiter.Commit(reserved, 0) + err := conn.limiter.Can(uint64(len(p))) + if err != nil { return 0, err } n, err := conn.PacketConn.WriteTo(p, addr) - conn.limiter.Commit(reserved, uint64(n)) - return n, err + if err != nil { + return 0, err + } + err = conn.limiter.Add(uint64(n)) + if err != nil { + return 0, err + } + return n, nil } type packetConnWithUploadTrafficLimiter struct { @@ -81,16 +94,19 @@ type packetConnWithUploadTrafficLimiter struct { } func (conn *packetConnWithUploadTrafficLimiter) ReadFrom(p []byte) (int, net.Addr, error) { - reserved, err := conn.limiter.Reserve(uint64(len(p))) - if reserved == 0 { + err := conn.limiter.Can(1) + if err != nil { return 0, nil, err } - if reserved < uint64(len(p)) { - p = p[:reserved] - } n, addr, err := conn.PacketConn.ReadFrom(p) - conn.limiter.Commit(reserved, uint64(n)) - return n, addr, err + if err != nil { + return n, nil, err + } + err = conn.limiter.Add(uint64(n)) + if err != nil { + return 0, nil, err + } + return n, addr, nil } func connWithDownloadTrafficWrapper(ctx context.Context, conn net.Conn, limiter TrafficLimiter, reverse bool) net.Conn { diff --git a/protocol/limiter/traffic/limiter.go b/protocol/limiter/traffic/limiter.go index d769701f..b1b95cbf 100644 --- a/protocol/limiter/traffic/limiter.go +++ b/protocol/limiter/traffic/limiter.go @@ -1,6 +1,6 @@ package traffic type TrafficLimiter interface { - Reserve(n uint64) (uint64, error) - Commit(reserved uint64, n uint64) + Can(n uint64) error + Add(n uint64) error } diff --git a/protocol/limiter/traffic/strategy.go b/protocol/limiter/traffic/strategy.go index 68a58b5d..d3ba5742 100644 --- a/protocol/limiter/traffic/strategy.go +++ b/protocol/limiter/traffic/strategy.go @@ -40,7 +40,7 @@ func (s *DefaultWrapStrategy) wrapConn(ctx context.Context, conn net.Conn, metad if err != nil { return nil, err } - _, err = limiter.Reserve(0) + err = limiter.Can(1) if err != nil { return nil, err } @@ -52,7 +52,7 @@ func (s *DefaultWrapStrategy) wrapPacketConn(ctx context.Context, conn net.Packe if err != nil { return nil, err } - _, err = limiter.Reserve(0) + err = limiter.Can(1) if err != nil { return nil, err } diff --git a/service/node/limiter/traffic.go b/service/node/limiter/traffic.go index 527bf861..4f6f678b 100644 --- a/service/node/limiter/traffic.go +++ b/service/node/limiter/traffic.go @@ -160,10 +160,9 @@ func (i *TrafficLimiterStrategyManager) DeleteTrafficLimiter(username string) { } type TrafficLimiter struct { - manager CM.NodeManager - limiter CM.TrafficLimiter - new uint64 - reserved uint64 + manager CM.NodeManager + limiter CM.TrafficLimiter + new uint64 mtx sync.Mutex } @@ -172,34 +171,34 @@ func NewTrafficLimiter(manager CM.NodeManager, limiter CM.TrafficLimiter) *Traff return &TrafficLimiter{manager: manager, limiter: limiter} } -func (l *TrafficLimiter) Reserve(n uint64) (uint64, error) { +func (l *TrafficLimiter) Can(n uint64) error { l.mtx.Lock() defer l.mtx.Unlock() - used := l.limiter.RawUsed + l.reserved - if used >= l.limiter.RawQuota { - return 0, E.New("traffic limit exceeded") + if l.limiter.RawUsed == l.limiter.RawQuota { + return E.New("traffic limit exceeded") } - remaining := l.limiter.RawQuota - used - if n > remaining { - l.reserved += remaining - return remaining, E.New("traffic limit exceeded") + if l.limiter.RawUsed+n > l.limiter.RawQuota { + l.new += l.limiter.RawQuota - l.limiter.RawUsed + l.limiter.RawUsed = l.limiter.RawQuota + return E.New("traffic limit exceeded") } - l.reserved += n - return n, nil + return nil } -func (l *TrafficLimiter) Commit(reserved uint64, n uint64) { - if reserved == 0 && n == 0 { - return - } +func (l *TrafficLimiter) Add(n uint64) error { l.mtx.Lock() defer l.mtx.Unlock() - if reserved > l.reserved { - reserved = l.reserved + if l.limiter.RawUsed == l.limiter.RawQuota { + return E.New("traffic limit exceeded") + } + if l.limiter.RawUsed+n > l.limiter.RawQuota { + l.new += l.limiter.RawQuota - l.limiter.RawUsed + l.limiter.RawUsed = l.limiter.RawQuota + return E.New("traffic limit exceeded") } - l.reserved -= reserved l.limiter.RawUsed += n l.new += n + return nil } func (l *TrafficLimiter) UpdateRemainingTraffic() error {