mirror of
https://github.com/shtorm-7/sing-box-extended.git
synced 2026-09-15 21:00:27 +00:00
limiter: revert traffic limiter to Can/Add interface
This commit is contained in:
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user