main
go 145 lines 2.77 KB
Raw
1 package portal
2
3 import (
4 "errors"
5 "io"
6 "net"
7 "sync"
8 "sync/atomic"
9 "time"
10
11 "golang.org/x/sync/errgroup"
12
13 "github.com/gosuda/portal-tunnel/v2/portal/policy"
14 )
15
16 type proxy struct {
17 activeConns atomic.Int64
18 tcpBytes atomic.Int64
19 tcpLoadMu sync.Mutex
20 tcpLoadAt time.Time
21 tcpLoadBytes int64
22 }
23
24 func (p *proxy) bridge(left, right net.Conn, identityKey string, bpsManager *policy.BPSManager) {
25 p.activeConns.Add(1)
26 defer p.activeConns.Add(-1)
27
28 defer left.Close()
29 defer right.Close()
30
31 throttled := bpsManager != nil && bpsManager.IdentityBPS(identityKey) > 0
32 var group errgroup.Group
33 group.Go(func() error {
34 err := p.copy(right, left, identityKey, bpsManager, throttled)
35 closeWrite(right)
36 return err
37 })
38 group.Go(func() error {
39 err := p.copy(left, right, identityKey, bpsManager, throttled)
40 closeWrite(left)
41 return err
42 })
43 _ = group.Wait()
44 }
45
46 func (p *proxy) activeConnectionCount() int64 {
47 return p.activeConns.Load()
48 }
49
50 func (p *proxy) currentTCPBPS(now time.Time) float64 {
51 totalTCPBytes := p.tcpBytes.Load()
52
53 p.tcpLoadMu.Lock()
54 defer p.tcpLoadMu.Unlock()
55
56 if p.tcpLoadAt.IsZero() {
57 p.tcpLoadAt = now
58 p.tcpLoadBytes = totalTCPBytes
59 return 0
60 }
61
62 if elapsed := now.Sub(p.tcpLoadAt); elapsed > 0 {
63 tcpTrafficBPS := float64(totalTCPBytes-p.tcpLoadBytes) / elapsed.Seconds()
64 p.tcpLoadAt = now
65 p.tcpLoadBytes = totalTCPBytes
66 return tcpTrafficBPS
67 }
68
69 return 0
70 }
71
72 func (p *proxy) copy(dst, src net.Conn, identityKey string, bpsManager *policy.BPSManager, throttled bool) error {
73 // fast path
74 if !throttled {
75 _, err := io.Copy(&countingConn{Conn: dst, bytes: &p.tcpBytes}, src)
76 return err
77 }
78
79 buf := make([]byte, 32*1024)
80 for {
81 nr, readErr := src.Read(buf)
82 if nr > 0 {
83 data := buf[:nr]
84 for len(data) > 0 {
85 chunkSize := len(data)
86 if bpsManager != nil {
87 chunkSize = bpsManager.ThrottleIdentityBPS(identityKey, chunkSize)
88 }
89
90 n, err := dst.Write(data[:chunkSize])
91 if n > 0 {
92 p.tcpBytes.Add(int64(n))
93 data = data[n:]
94 }
95 if err != nil {
96 return err
97 }
98 if n == 0 {
99 return io.ErrShortWrite
100 }
101 }
102 }
103 if readErr != nil {
104 if errors.Is(readErr, io.EOF) {
105 return nil
106 }
107 return readErr
108 }
109 }
110 }
111
112 type countingConn struct {
113 net.Conn
114 bytes *atomic.Int64
115 }
116
117 func (c *countingConn) Write(p []byte) (int, error) {
118 n, err := c.Conn.Write(p)
119 if n > 0 {
120 c.bytes.Add(int64(n))
121 }
122 return n, err
123 }
124
125 func (c *countingConn) ReadFrom(r io.Reader) (int64, error) {
126 readerFrom, ok := c.Conn.(io.ReaderFrom)
127 if !ok {
128 return io.Copy(struct{ io.Writer }{Writer: c}, r)
129 }
130
131 n, err := readerFrom.ReadFrom(r)
132 if n > 0 {
133 c.bytes.Add(n)
134 }
135 return n, err
136 }
137
138 func closeWrite(conn net.Conn) {
139 type closeWriter interface {
140 CloseWrite() error
141 }
142 if cw, ok := conn.(closeWriter); ok {
143 _ = cw.CloseWrite()
144 }
145 }