main
go 397 lines 9.16 KB
Raw
1 package sdk
2
3 import (
4 "bufio"
5 "bytes"
6 "context"
7 "crypto/rand"
8 "crypto/tls"
9 "encoding/hex"
10 "errors"
11 "fmt"
12 "io"
13 "net"
14 "net/url"
15 "sync"
16 "time"
17
18 "github.com/rs/zerolog/log"
19
20 "github.com/gosuda/portal-tunnel/v2/portal/keyless"
21 "github.com/gosuda/portal-tunnel/v2/types"
22 "github.com/gosuda/portal-tunnel/v2/utils"
23 )
24
25 const (
26 mitmProbeExporterLabel = "Portal-MITM-Probe-v1"
27 mitmProbePeekTimeout = 100 * time.Millisecond
28 mitmProbePaddingMin = 96
29 mitmProbePaddingMax = 320
30
31 defaultMITMProbeCooldown = 30 * time.Second
32 defaultMITMProbeTimeout = 5 * time.Second
33 )
34
35 type MITMProbeReport struct {
36 RelayURL string
37 PublicURL string
38 Address string
39 ECHAccepted bool
40 Detected bool
41 Reason string
42 }
43
44 type mitmProbePending struct {
45 expected []byte
46 resultCh chan string
47 }
48
49 type mitmManager struct {
50 ctx context.Context
51 listener *listener
52 ban bool
53
54 mu sync.Mutex
55 pending map[string]*mitmProbePending
56 inFlight bool
57 lastAt time.Time
58 }
59
60 func newMITMManager(ctx context.Context, listener *listener, ban bool) *mitmManager {
61 return &mitmManager{
62 ctx: ctx,
63 ban: ban,
64 listener: listener,
65 pending: make(map[string]*mitmProbePending),
66 }
67 }
68
69 func (m *mitmManager) reset() {
70 m.mu.Lock()
71 clear(m.pending)
72 m.inFlight = false
73 m.lastAt = time.Time{}
74 m.mu.Unlock()
75 }
76
77 func (m *mitmManager) probeTLSPassthrough(ctx context.Context) (MITMProbeReport, error) {
78 l := m.listener
79 if l == nil || l.relayURL == nil {
80 return MITMProbeReport{}, errors.New("listener is not ready")
81 }
82
83 lease, ok := l.leaseSnapshot()
84 if !ok {
85 return MITMProbeReport{}, errors.New("listener is not registered")
86 }
87 if lease.hostname == "" {
88 return MITMProbeReport{}, errors.New("listener hostname is unavailable")
89 }
90
91 publicURL := l.publicURLForLease(lease)
92 if publicURL == "" {
93 return MITMProbeReport{}, errors.New("listener is not registered")
94 }
95
96 report := MITMProbeReport{
97 RelayURL: l.relayURL.String(),
98 PublicURL: publicURL,
99 Address: l.identity.Address,
100 }
101
102 probeCtx, cancel := context.WithTimeout(ctx, defaultMITMProbeTimeout)
103 defer cancel()
104
105 nonceRaw := make([]byte, 16)
106 if _, err := io.ReadFull(rand.Reader, nonceRaw); err != nil {
107 return report, fmt.Errorf("generate probe nonce: %w", err)
108 }
109 nonceHex := hex.EncodeToString(nonceRaw)
110
111 dialAddr, err := m.probeDialAddress(publicURL)
112 if err != nil {
113 return report, err
114 }
115
116 probeTLSConf := &tls.Config{
117 ServerName: lease.hostname,
118 InsecureSkipVerify: true,
119 MinVersion: keyless.MinTLSVersion(len(lease.echConfigList) > 0),
120 EncryptedClientHelloConfigList: bytes.Clone(lease.echConfigList),
121 }
122
123 dialer := &tls.Dialer{
124 NetDialer: &net.Dialer{Timeout: l.dialTimeout},
125 Config: probeTLSConf,
126 }
127 conn, err := dialer.DialContext(probeCtx, "tcp", dialAddr)
128 if err != nil {
129 return report, fmt.Errorf("dial mitm probe: %w", err)
130 }
131 defer conn.Close()
132
133 tlsConn, ok := conn.(*tls.Conn)
134 if !ok {
135 return report, errors.New("mitm probe connection is not tls")
136 }
137
138 clientState := tlsConn.ConnectionState()
139 report.ECHAccepted = clientState.ECHAccepted
140 expected, err := (&clientState).ExportKeyingMaterial(mitmProbeExporterLabel, nil, 32)
141 if err != nil {
142 return report, fmt.Errorf("export client probe keying material: %w", err)
143 }
144 resultCh, cleanupProbe := m.startProbe(nonceHex, expected)
145 defer cleanupProbe()
146
147 paddingLen := mitmProbePaddingMin
148 if paddingRange := mitmProbePaddingMax - mitmProbePaddingMin; paddingRange > 0 {
149 var paddingSeed [1]byte
150 if _, err := io.ReadFull(rand.Reader, paddingSeed[:]); err != nil {
151 return report, fmt.Errorf("generate probe padding length: %w", err)
152 }
153 paddingLen += int(paddingSeed[0]) % (paddingRange + 1)
154 }
155
156 frame := make([]byte, len(nonceRaw)+paddingLen)
157 if _, err := io.ReadFull(rand.Reader, frame); err != nil {
158 return report, fmt.Errorf("generate probe frame: %w", err)
159 }
160 copy(frame, nonceRaw)
161 if _, err := conn.Write(frame); err != nil {
162 return report, fmt.Errorf("write mitm probe: %w", err)
163 }
164
165 select {
166 case reason := <-resultCh:
167 report.Detected = reason != ""
168 report.Reason = reason
169 case <-probeCtx.Done():
170 report.Reason = types.MITMProbeReasonProbeTimeout
171 if !errors.Is(probeCtx.Err(), context.DeadlineExceeded) {
172 return report, probeCtx.Err()
173 }
174 }
175 return report, nil
176 }
177
178 func (m *mitmManager) probeDialAddress(publicURL string) (string, error) {
179 l := m.listener
180 if l == nil || l.relayURL == nil {
181 return "", errors.New("listener is not ready")
182 }
183 parsedURL, err := url.Parse(publicURL)
184 if err != nil {
185 return "", fmt.Errorf("parse public url: %w", err)
186 }
187
188 dialHost := parsedURL.Host
189 if utils.IsLocalRelayHost(l.relayURL.Hostname()) {
190 dialHost = l.relayURL.Host
191 }
192 return utils.EnsurePort(dialHost), nil
193 }
194
195 func (m *mitmManager) maybeStart() {
196 l := m.listener
197 select {
198 case <-l.doneCh:
199 return
200 default:
201 }
202
203 m.mu.Lock()
204 if m.inFlight || !m.lastAt.IsZero() && time.Since(m.lastAt) < defaultMITMProbeCooldown {
205 m.mu.Unlock()
206 return
207 }
208 m.inFlight = true
209 m.mu.Unlock()
210
211 go func() {
212 report, err := m.probeTLSPassthrough(m.ctx)
213 success := err == nil && report.Reason != types.MITMProbeReasonProbeTimeout
214 m.mu.Lock()
215 m.inFlight = false
216 if success {
217 m.lastAt = time.Now()
218 }
219 m.mu.Unlock()
220 m.logResult(report, err)
221 }()
222 }
223
224 func (m *mitmManager) logResult(report MITMProbeReport, err error) {
225 l := m.listener
226 if l == nil {
227 return
228 }
229 closed := false
230 select {
231 case <-l.doneCh:
232 closed = true
233 default:
234 }
235 relayURL := ""
236 if l.relayURL != nil {
237 relayURL = l.relayURL.String()
238 }
239 switch {
240 case closed:
241 return
242 case err != nil:
243 if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
244 return
245 }
246 log.Warn().
247 Err(err).
248 Bool("ech_accepted", report.ECHAccepted).
249 Str("relay_url", relayURL).
250 Str("address", l.identity.Address).
251 Msg("tls passthrough self-probe failed")
252 case report.Reason == types.MITMProbeReasonProbeTimeout:
253 log.Warn().
254 Bool("ech_accepted", report.ECHAccepted).
255 Str("relay_url", report.RelayURL).
256 Str("public_url", report.PublicURL).
257 Str("address", report.Address).
258 Msg("tls self-probe timed out before passthrough could be verified")
259 case report.Detected:
260 event := log.Warn().
261 Bool("ban_mitm", m.ban).
262 Bool("ech_accepted", report.ECHAccepted).
263 Str("reason", report.Reason).
264 Str("relay_url", report.RelayURL).
265 Str("public_url", report.PublicURL).
266 Str("address", report.Address)
267 if m.ban {
268 event.Msg("tls termination suspected by self-probe; banning relay")
269 if l.relaySet != nil && report.RelayURL != "" {
270 l.relaySet.UnconfirmRelayURL(report.RelayURL)
271 l.relaySet.BanRelayURL(report.RelayURL)
272 }
273 _ = l.Close()
274 return
275 }
276 event.Msg("tls termination suspected by self-probe")
277 default:
278 log.Debug().
279 Bool("ech_accepted", report.ECHAccepted).
280 Str("relay_url", report.RelayURL).
281 Str("public_url", report.PublicURL).
282 Str("address", report.Address).
283 Msg("tls passthrough self-probe passed")
284 }
285 }
286
287 func (m *mitmManager) maybeHandleConn(conn net.Conn) (net.Conn, bool, error) {
288 if conn == nil {
289 return conn, false, nil
290 }
291
292 m.mu.Lock()
293 hasPending := len(m.pending) > 0
294 m.mu.Unlock()
295 if !hasPending {
296 return conn, false, nil
297 }
298
299 tlsConn, ok := conn.(*tls.Conn)
300 if !ok {
301 return conn, false, nil
302 }
303
304 frameSize := 16
305 reader := bufio.NewReaderSize(conn, frameSize)
306 _ = conn.SetReadDeadline(time.Now().Add(mitmProbePeekTimeout))
307 peeked, err := reader.Peek(frameSize)
308 defer conn.SetReadDeadline(time.Time{})
309 if err != nil {
310 return wrapBufferedConn(conn, reader), false, nil
311 }
312
313 nonceHex := hex.EncodeToString(peeked[:frameSize])
314 m.mu.Lock()
315 _, ok = m.pending[nonceHex]
316 m.mu.Unlock()
317 if !ok {
318 return wrapBufferedConn(conn, reader), false, nil
319 }
320
321 defer conn.Close()
322
323 frame := make([]byte, frameSize)
324 if _, err := io.ReadFull(reader, frame); err != nil {
325 return nil, true, fmt.Errorf("read mitm probe frame: %w", err)
326 }
327
328 serverState := tlsConn.ConnectionState()
329 actual, err := (&serverState).ExportKeyingMaterial(mitmProbeExporterLabel, nil, 32)
330 if err != nil {
331 return nil, true, fmt.Errorf("export server probe keying material: %w", err)
332 }
333
334 m.completeProbe(nonceHex, actual)
335 return nil, true, nil
336 }
337
338 func (m *mitmManager) startProbe(nonce string, expected []byte) (<-chan string, func()) {
339 m.mu.Lock()
340 state := &mitmProbePending{
341 expected: bytes.Clone(expected),
342 resultCh: make(chan string, 1),
343 }
344 m.pending[nonce] = state
345 m.mu.Unlock()
346
347 return state.resultCh, func() {
348 m.mu.Lock()
349 delete(m.pending, nonce)
350 m.mu.Unlock()
351 }
352 }
353
354 func (m *mitmManager) completeProbe(nonce string, actual []byte) {
355 m.mu.Lock()
356 state := m.pending[nonce]
357 m.mu.Unlock()
358 if state == nil {
359 return
360 }
361
362 reason := ""
363 if !bytes.Equal(state.expected, actual) {
364 reason = types.MITMProbeReasonExporterMismatch
365 }
366
367 select {
368 case state.resultCh <- reason:
369 default:
370 }
371 }
372
373 type mitmProbeConn struct {
374 net.Conn
375 manager *mitmManager
376 startOnce sync.Once
377 }
378
379 func (c *mitmProbeConn) Read(p []byte) (int, error) {
380 n, err := c.Conn.Read(p)
381 if n > 0 {
382 c.startOnce.Do(func() {
383 c.manager.maybeStart()
384 })
385 }
386 return n, err
387 }
388
389 func (c *mitmProbeConn) Write(p []byte) (int, error) {
390 n, err := c.Conn.Write(p)
391 if n > 0 {
392 c.startOnce.Do(func() {
393 c.manager.maybeStart()
394 })
395 }
396 return n, err
397 }