| 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 | } |