limiter: revert traffic limiter to Can/Add interface

This commit is contained in:
Shtorm
2026-09-04 01:48:30 +03:00
parent 673a408dca
commit 301df12dee
4 changed files with 64 additions and 49 deletions
+40 -24
View File
@@ -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 {
+2 -2
View File
@@ -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
}
+2 -2
View File
@@ -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
}
+20 -21
View File
@@ -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 {