main
go 395 lines 10.6 KB
Raw
1 package sdk
2
3 import (
4 "bytes"
5 "context"
6 "crypto/ecdsa"
7 "crypto/elliptic"
8 "crypto/rand"
9 "crypto/tls"
10 "crypto/x509"
11 "crypto/x509/pkix"
12 "encoding/hex"
13 "encoding/pem"
14 "io"
15 "math/big"
16 "net"
17 "net/url"
18 "testing"
19 "time"
20
21 "github.com/gosuda/portal-tunnel/v2/portal/discovery"
22 "github.com/gosuda/portal-tunnel/v2/types"
23 )
24
25 func TestMITMProbeConnMatchesExporter(t *testing.T) {
26 clientConn, serverConn := newMITMProbeTLSPair(t)
27 defer closeMITMProbeTLSConn(clientConn)
28 defer closeMITMProbeTLSConn(serverConn)
29
30 listener := &listener{}
31 listener.mitmManager = newMITMManager(context.Background(), listener, false)
32
33 nonce := make([]byte, 16)
34 if _, err := rand.Read(nonce); err != nil {
35 t.Fatalf("rand.Read() error = %v", err)
36 }
37 nonceHex := hex.EncodeToString(nonce)
38 clientState := clientConn.ConnectionState()
39 expected, err := (&clientState).ExportKeyingMaterial(mitmProbeExporterLabel, nil, 32)
40 if err != nil {
41 t.Fatalf("client ExportKeyingMaterial() error = %v", err)
42 }
43 resultCh, cleanupProbe := listener.mitmManager.startProbe(nonceHex, expected)
44 defer cleanupProbe()
45
46 handleDone := make(chan struct{})
47 go func() {
48 defer close(handleDone)
49 nextConn, handled, err := listener.mitmManager.maybeHandleConn(serverConn)
50 if err != nil {
51 t.Errorf("maybeHandleConn() error = %v", err)
52 return
53 }
54 if nextConn != nil {
55 t.Error("maybeHandleConn() returned passthrough conn for probe")
56 }
57 if !handled {
58 t.Error("maybeHandleConn() handled = false, want true")
59 }
60 }()
61
62 frame := bytes.Clone(nonce)
63 frame = append(frame, bytes.Repeat([]byte{0xAB}, 128)...)
64 if _, err := clientConn.Write(frame); err != nil {
65 t.Fatalf("clientConn.Write() error = %v", err)
66 }
67 _ = clientConn.Close()
68
69 select {
70 case reason := <-resultCh:
71 if reason != "" {
72 t.Fatalf("probe reason = %q, want empty", reason)
73 }
74 case <-time.After(2 * time.Second):
75 t.Fatal("timed out waiting for probe result")
76 }
77
78 select {
79 case <-handleDone:
80 case <-time.After(2 * time.Second):
81 t.Fatal("timed out waiting for probe handler")
82 }
83 }
84
85 func TestMITMProbeConnDetectsExporterMismatch(t *testing.T) {
86 clientConn, serverConn := newMITMProbeTLSPair(t)
87 defer closeMITMProbeTLSConn(clientConn)
88 defer closeMITMProbeTLSConn(serverConn)
89
90 listener := &listener{}
91 listener.mitmManager = newMITMManager(context.Background(), listener, false)
92
93 nonce := make([]byte, 16)
94 if _, err := rand.Read(nonce); err != nil {
95 t.Fatalf("rand.Read() error = %v", err)
96 }
97 nonceHex := hex.EncodeToString(nonce)
98 resultCh, cleanupProbe := listener.mitmManager.startProbe(nonceHex, make([]byte, 32))
99 defer cleanupProbe()
100
101 handleDone := make(chan struct{})
102 go func() {
103 defer close(handleDone)
104 nextConn, handled, err := listener.mitmManager.maybeHandleConn(serverConn)
105 if err != nil {
106 t.Errorf("maybeHandleConn() error = %v", err)
107 return
108 }
109 if nextConn != nil {
110 t.Error("maybeHandleConn() returned passthrough conn for probe")
111 }
112 if !handled {
113 t.Error("maybeHandleConn() handled = false, want true")
114 }
115 }()
116
117 frame := bytes.Clone(nonce)
118 frame = append(frame, bytes.Repeat([]byte{0xCD}, 128)...)
119 if _, err := clientConn.Write(frame); err != nil {
120 t.Fatalf("clientConn.Write() error = %v", err)
121 }
122 _ = clientConn.Close()
123
124 select {
125 case reason := <-resultCh:
126 if reason != types.MITMProbeReasonExporterMismatch {
127 t.Fatalf("probe reason = %q, want %q", reason, types.MITMProbeReasonExporterMismatch)
128 }
129 case <-time.After(2 * time.Second):
130 t.Fatal("timed out waiting for probe result")
131 }
132
133 select {
134 case <-handleDone:
135 case <-time.After(2 * time.Second):
136 t.Fatal("timed out waiting for probe handler")
137 }
138 }
139
140 func TestMITMProbeConnPassesThroughNormalTraffic(t *testing.T) {
141 clientConn, serverConn := newMITMProbeTLSPair(t)
142 defer closeMITMProbeTLSConn(clientConn)
143 defer closeMITMProbeTLSConn(serverConn)
144
145 listener := &listener{}
146 listener.mitmManager = newMITMManager(context.Background(), listener, false)
147
148 type handleResult struct {
149 conn net.Conn
150 handled bool
151 err error
152 }
153 handleResultCh := make(chan handleResult, 1)
154 go func() {
155 nextConn, handled, err := listener.mitmManager.maybeHandleConn(serverConn)
156 handleResultCh <- handleResult{conn: nextConn, handled: handled, err: err}
157 }()
158
159 payload := []byte("GET / HTTP/1.1\r\nHost: localhost\r\n\r\n")
160 var result handleResult
161 select {
162 case result = <-handleResultCh:
163 case <-time.After(2 * time.Second):
164 t.Fatal("timed out waiting for passthrough result")
165 }
166
167 if result.err != nil {
168 t.Fatalf("maybeHandleConn() error = %v", result.err)
169 }
170 if result.handled {
171 t.Fatal("maybeHandleConn() handled = true, want false")
172 }
173 if result.conn == nil {
174 t.Fatal("maybeHandleConn() returned nil passthrough conn")
175 }
176
177 writeErrCh := make(chan error, 1)
178 go func() {
179 _, err := clientConn.Write(payload)
180 writeErrCh <- err
181 }()
182
183 got := make([]byte, len(payload))
184 if _, err := io.ReadFull(result.conn, got); err != nil {
185 t.Fatalf("ReadFull() error = %v", err)
186 }
187 select {
188 case err := <-writeErrCh:
189 if err != nil {
190 t.Fatalf("clientConn.Write() error = %v", err)
191 }
192 case <-time.After(2 * time.Second):
193 t.Fatal("timed out waiting for client write")
194 }
195 if !bytes.Equal(got, payload) {
196 t.Fatalf("passthrough payload = %q, want %q", got, payload)
197 }
198 }
199
200 func TestMITMProbeDetectionBansListener(t *testing.T) {
201 doneCh := make(chan struct{})
202 relayURL, err := url.Parse("https://relay.example")
203 if err != nil {
204 t.Fatalf("url.Parse() error = %v", err)
205 }
206
207 listener := &listener{
208 relayURL: relayURL,
209 relaySet: mustRelaySet(t, relayURL.String()),
210 cancel: func() {
211 select {
212 case <-doneCh:
213 default:
214 close(doneCh)
215 }
216 },
217 doneCh: doneCh,
218 }
219 listener.mitmManager = newMITMManager(context.Background(), listener, true)
220
221 listener.mitmManager.logResult(MITMProbeReport{
222 RelayURL: relayURL.String(),
223 Detected: true,
224 Reason: types.MITMProbeReasonExporterMismatch,
225 }, nil)
226
227 routes, err := listener.relaySet.PlanRoutes(nil, discovery.RouteState{})
228 if err != nil {
229 t.Fatalf("PlanRoutes() error = %v", err)
230 }
231 for _, route := range routes {
232 if route.ListenerRelayURL() == relayURL.String() {
233 t.Fatal("relay still active after mitm detection")
234 }
235 }
236 select {
237 case <-listener.doneCh:
238 default:
239 t.Fatal("listener.doneCh is open, want closed")
240 }
241 }
242
243 func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
244 doneCh := make(chan struct{})
245 relayURL, err := url.Parse("https://relay.example")
246 if err != nil {
247 t.Fatalf("url.Parse() error = %v", err)
248 }
249
250 listener := &listener{
251 relayURL: relayURL,
252 relaySet: mustRelaySet(t, relayURL.String()),
253 doneCh: doneCh,
254 }
255 listener.mitmManager = newMITMManager(context.Background(), listener, false)
256
257 listener.mitmManager.logResult(MITMProbeReport{
258 RelayURL: relayURL.String(),
259 Detected: true,
260 Reason: types.MITMProbeReasonExporterMismatch,
261 }, nil)
262
263 routes, err := listener.relaySet.PlanRoutes(nil, discovery.RouteState{})
264 if err != nil {
265 t.Fatalf("PlanRoutes() error = %v", err)
266 }
267 activeRelayURLs := make([]string, 0, len(routes))
268 for _, route := range routes {
269 activeRelayURLs = append(activeRelayURLs, route.ListenerRelayURL())
270 }
271 if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {
272 t.Fatalf("ActiveRelayURLs() = %v, want [%q]", activeRelayURLs, relayURL.String())
273 }
274 select {
275 case <-listener.doneCh:
276 t.Fatal("listener.doneCh is closed, want open")
277 default:
278 }
279 }
280
281 func TestMITMProbeDialAddressUsesRelayHostForLocalRelay(t *testing.T) {
282 relayURL, err := url.Parse("https://localhost:4017")
283 if err != nil {
284 t.Fatalf("url.Parse() error = %v", err)
285 }
286
287 listener := &listener{
288 relayURL: relayURL,
289 }
290 listener.mitmManager = newMITMManager(context.Background(), listener, false)
291
292 got, err := listener.mitmManager.probeDialAddress("https://bravo-gecko-disco.localhost:4017")
293 if err != nil {
294 t.Fatalf("probeDialAddress() error = %v", err)
295 }
296 if got != "localhost:4017" {
297 t.Fatalf("probeDialAddress() = %q, want %q", got, "localhost:4017")
298 }
299 }
300
301 func TestMITMProbeDialAddressUsesPublicURLForRemoteRelay(t *testing.T) {
302 relayURL, err := url.Parse("https://relay.example")
303 if err != nil {
304 t.Fatalf("url.Parse() error = %v", err)
305 }
306
307 listener := &listener{
308 relayURL: relayURL,
309 }
310 listener.mitmManager = newMITMManager(context.Background(), listener, false)
311
312 got, err := listener.mitmManager.probeDialAddress("https://bravo-gecko-disco.example")
313 if err != nil {
314 t.Fatalf("probeDialAddress() error = %v", err)
315 }
316 if got != "bravo-gecko-disco.example:443" {
317 t.Fatalf("probeDialAddress() = %q, want %q", got, "bravo-gecko-disco.example:443")
318 }
319 }
320
321 func newMITMProbeTLSPair(t *testing.T) (*tls.Conn, *tls.Conn) {
322 t.Helper()
323
324 cert := newMITMProbeCertificate(t)
325 clientRaw, serverRaw := net.Pipe()
326 clientConn := tls.Client(clientRaw, &tls.Config{
327 InsecureSkipVerify: true,
328 MinVersion: tls.VersionTLS13,
329 NextProtos: []string{"http/1.1"},
330 })
331 serverConn := tls.Server(serverRaw, &tls.Config{
332 Certificates: []tls.Certificate{cert},
333 MinVersion: tls.VersionTLS13,
334 NextProtos: []string{"http/1.1"},
335 })
336
337 errCh := make(chan error, 2)
338 go func() { errCh <- serverConn.HandshakeContext(context.Background()) }()
339 go func() { errCh <- clientConn.HandshakeContext(context.Background()) }()
340 for range 2 {
341 if err := <-errCh; err != nil {
342 t.Fatalf("TLS handshake error = %v", err)
343 }
344 }
345
346 return clientConn, serverConn
347 }
348
349 func closeMITMProbeTLSConn(conn *tls.Conn) {
350 if conn == nil {
351 return
352 }
353 _ = conn.SetDeadline(time.Now())
354 _ = conn.Close()
355 }
356
357 func newMITMProbeCertificate(t *testing.T) tls.Certificate {
358 t.Helper()
359
360 privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
361 if err != nil {
362 t.Fatalf("GenerateKey() error = %v", err)
363 }
364
365 template := &x509.Certificate{
366 SerialNumber: big.NewInt(1),
367 Subject: pkix.Name{
368 CommonName: "portal-mitm-probe",
369 },
370 NotBefore: time.Now().Add(-time.Hour),
371 NotAfter: time.Now().Add(time.Hour),
372 KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
373 ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
374 BasicConstraintsValid: true,
375 DNSNames: []string{"localhost"},
376 }
377
378 der, err := x509.CreateCertificate(rand.Reader, template, template, &privateKey.PublicKey, privateKey)
379 if err != nil {
380 t.Fatalf("CreateCertificate() error = %v", err)
381 }
382
383 keyDER, err := x509.MarshalPKCS8PrivateKey(privateKey)
384 if err != nil {
385 t.Fatalf("MarshalPKCS8PrivateKey() error = %v", err)
386 }
387
388 certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
389 keyPEM := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyDER})
390 cert, err := tls.X509KeyPair(certPEM, keyPEM)
391 if err != nil {
392 t.Fatalf("X509KeyPair() error = %v", err)
393 }
394 return cert
395 }