main
go 349 lines 6.48 KB
Raw
1 package transport
2
3 import (
4 "context"
5 "errors"
6 "fmt"
7 "net"
8 "sync"
9 "time"
10
11 "github.com/rs/zerolog/log"
12
13 "github.com/gosuda/portal-tunnel/v2/types"
14 )
15
16 const defaultSessionWriteLimit = 5 * time.Second
17
18 var errStreamFull = errors.New("stream ready queue full")
19
20 type RelayStream struct {
21 notify chan struct{}
22 identityKey string
23 ready []*relaySession
24 idleInterval time.Duration
25 readyLimit int
26 closedErr error
27 mu sync.Mutex
28 }
29
30 func NewRelayStream(identityKey string, idleInterval time.Duration, readyLimit int) *RelayStream {
31 return &RelayStream{
32 identityKey: identityKey,
33 idleInterval: idleInterval,
34 readyLimit: readyLimit,
35 notify: make(chan struct{}, 1),
36 }
37 }
38
39 func (b *RelayStream) OfferConn(conn net.Conn) error {
40 if conn == nil {
41 return errors.New("reverse connection is required")
42 }
43 session := newRelaySession(conn, b.idleInterval)
44
45 b.mu.Lock()
46 if b.closedErr != nil {
47 err := b.closedErr
48 b.mu.Unlock()
49 _ = session.Close()
50 return err
51 }
52
53 if b.readyLimit > 0 && len(b.ready) >= b.readyLimit {
54 b.mu.Unlock()
55 _ = session.Close()
56 return errStreamFull
57 }
58
59 session.StartIdle()
60 b.ready = append(b.ready, session)
61 b.signalLocked()
62 b.mu.Unlock()
63
64 go b.watchSession(session)
65 return nil
66 }
67
68 func (b *RelayStream) Claim(ctx context.Context) (net.Conn, error) {
69 return b.claimWithMarker(ctx, types.MarkerTLSStart)
70 }
71
72 func (b *RelayStream) claimRaw(ctx context.Context) (net.Conn, error) {
73 return b.claimWithMarker(ctx, types.MarkerRawStart)
74 }
75
76 func (b *RelayStream) claimWithMarker(ctx context.Context, marker byte) (net.Conn, error) {
77 for {
78 b.mu.Lock()
79 if b.closedErr != nil {
80 err := b.closedErr
81 b.mu.Unlock()
82 return nil, err
83 }
84
85 if len(b.ready) > 0 {
86 session := b.ready[0]
87 b.ready = b.ready[1:]
88 b.mu.Unlock()
89
90 if session.IsClosed() {
91 continue
92 }
93 if err := session.activateWithMarker(marker); err != nil {
94 _ = session.Close()
95 continue
96 }
97 return session, nil
98 }
99 b.mu.Unlock()
100
101 select {
102 case <-ctx.Done():
103 return nil, ctx.Err()
104 case <-b.notify:
105 }
106 }
107 }
108
109 func (b *RelayStream) Close() {
110 b.mu.Lock()
111 sessions := b.ready
112 b.ready = nil
113 if b.closedErr == nil {
114 b.closedErr = net.ErrClosed
115 }
116 b.signalLocked()
117 b.mu.Unlock()
118
119 for _, session := range sessions {
120 _ = session.Close()
121 }
122 }
123
124 func (b *RelayStream) ReadyCount() int {
125 b.mu.Lock()
126 defer b.mu.Unlock()
127 return len(b.ready)
128 }
129
130 func (b *RelayStream) watchSession(session *relaySession) {
131 <-session.Done()
132
133 var readyCount int
134
135 b.mu.Lock()
136 for i := range b.ready {
137 if b.ready[i] == session {
138 b.ready = append(b.ready[:i], b.ready[i+1:]...)
139 break
140 }
141 }
142 readyCount = len(b.ready)
143 log.Info().
144 Str("identity_key", b.identityKey).
145 Str("remote_addr", session.remoteAddrString()).
146 Int("ready", readyCount).
147 Msg("sdk reverse disconnected")
148 b.signalLocked()
149 b.mu.Unlock()
150 }
151
152 func (b *RelayStream) signalLocked() {
153 select {
154 case b.notify <- struct{}{}:
155 default:
156 }
157 }
158
159 type sessionState int
160
161 const (
162 sessionIdle sessionState = iota
163 sessionClaimed
164 sessionClosed
165 )
166
167 type relaySession struct {
168 conn net.Conn
169 keepaliveStop chan struct{}
170 keepaliveDone chan struct{}
171 done chan struct{}
172 idleInterval time.Duration
173 state sessionState
174 closeOnce sync.Once
175 mu sync.Mutex
176 }
177
178 func newRelaySession(conn net.Conn, idleInterval time.Duration) *relaySession {
179 return &relaySession{
180 conn: conn,
181 idleInterval: idleInterval,
182 state: sessionIdle,
183 done: make(chan struct{}),
184 }
185 }
186
187 func (s *relaySession) Read(p []byte) (int, error) {
188 return s.conn.Read(p)
189 }
190
191 func (s *relaySession) Write(p []byte) (int, error) {
192 return s.conn.Write(p)
193 }
194
195 func (s *relaySession) LocalAddr() net.Addr {
196 if s == nil || s.conn == nil {
197 return nil
198 }
199 return s.conn.LocalAddr()
200 }
201
202 func (s *relaySession) RemoteAddr() net.Addr {
203 if s == nil || s.conn == nil {
204 return nil
205 }
206 return s.conn.RemoteAddr()
207 }
208
209 func (s *relaySession) SetDeadline(t time.Time) error {
210 return s.conn.SetDeadline(t)
211 }
212
213 func (s *relaySession) SetReadDeadline(t time.Time) error {
214 return s.conn.SetReadDeadline(t)
215 }
216
217 func (s *relaySession) SetWriteDeadline(t time.Time) error {
218 return s.conn.SetWriteDeadline(t)
219 }
220
221 func (s *relaySession) Done() <-chan struct{} {
222 return s.done
223 }
224
225 func (s *relaySession) remoteAddrString() string {
226 if s == nil || s.conn == nil || s.conn.RemoteAddr() == nil {
227 return ""
228 }
229 return s.conn.RemoteAddr().String()
230 }
231
232 func (s *relaySession) IsClosed() bool {
233 select {
234 case <-s.done:
235 return true
236 default:
237 return false
238 }
239 }
240
241 func (s *relaySession) StartIdle() {
242 s.mu.Lock()
243 if s.state != sessionIdle || s.keepaliveStop != nil {
244 s.mu.Unlock()
245 return
246 }
247 stop := make(chan struct{})
248 done := make(chan struct{})
249 s.keepaliveStop = stop
250 s.keepaliveDone = done
251 s.mu.Unlock()
252
253 go s.runKeepalive(stop, done)
254 }
255
256 func (s *relaySession) Activate() error {
257 return s.activateWithMarker(types.MarkerTLSStart)
258 }
259
260 func (s *relaySession) activateWithMarker(marker byte) error {
261 s.mu.Lock()
262 if s.state != sessionIdle {
263 state := s.state
264 s.mu.Unlock()
265 return fmt.Errorf("session not idle: %d", state)
266 }
267 stop := s.keepaliveStop
268 done := s.keepaliveDone
269 s.keepaliveStop = nil
270 s.keepaliveDone = nil
271 s.state = sessionClaimed
272 s.mu.Unlock()
273
274 if stop != nil {
275 close(stop)
276 }
277 if done != nil {
278 <-done
279 }
280
281 s.mu.Lock()
282 defer s.mu.Unlock()
283 if s.state == sessionClosed {
284 return net.ErrClosed
285 }
286 _ = s.conn.SetWriteDeadline(time.Now().Add(defaultSessionWriteLimit))
287 _, err := s.conn.Write([]byte{marker})
288 _ = s.conn.SetWriteDeadline(time.Time{})
289 if err != nil {
290 _ = s.Close()
291 }
292 return err
293 }
294
295 func (s *relaySession) Close() error {
296 var err error
297 s.closeOnce.Do(func() {
298 s.mu.Lock()
299 stop := s.keepaliveStop
300 done := s.keepaliveDone
301 s.keepaliveStop = nil
302 s.keepaliveDone = nil
303 s.state = sessionClosed
304 conn := s.conn
305 s.mu.Unlock()
306
307 if stop != nil {
308 close(stop)
309 }
310 if done != nil {
311 <-done
312 }
313
314 err = conn.Close()
315 close(s.done)
316 })
317 return err
318 }
319
320 func (s *relaySession) runKeepalive(stop <-chan struct{}, done chan<- struct{}) {
321 defer close(done)
322
323 ticker := time.NewTicker(s.idleInterval)
324 defer ticker.Stop()
325
326 for {
327 select {
328 case <-stop:
329 return
330 case <-s.done:
331 return
332 case <-ticker.C:
333 }
334
335 s.mu.Lock()
336 if s.state != sessionIdle {
337 s.mu.Unlock()
338 return
339 }
340 _ = s.conn.SetWriteDeadline(time.Now().Add(defaultSessionWriteLimit))
341 _, err := s.conn.Write([]byte{types.MarkerKeepalive})
342 _ = s.conn.SetWriteDeadline(time.Time{})
343 s.mu.Unlock()
344 if err != nil {
345 _ = s.Close()
346 return
347 }
348 }
349 }