package main import ( "crypto/tls" "crypto/x509" "io" "net" "testing" ) func TestCertMatchesFingerprint(t *testing.T) { certDER, _, err := generateHubCert() if err != nil { t.Fatal(err) } fp := certFingerprint(certDER) if !certMatchesFingerprint(certDER, fp) { t.Fatal("should match its own fingerprint") } if !certMatchesFingerprint(certDER, " "+fp+" ") { t.Fatal("should tolerate surrounding whitespace") } if certMatchesFingerprint(certDER, "00"+fp[2:]) { t.Fatal("should not match a different fingerprint") } } func TestValidFingerprint(t *testing.T) { certDER, _, err := generateHubCert() if err != nil { t.Fatal(err) } fp := certFingerprint(certDER) cases := []struct { s string ok bool }{ {fp, true}, {"", false}, {"not-hex-at-all-not-hex-at-all-not-hex-at-all-not-hex-at-all-00", false}, {fp[:len(fp)-1], false}, // one char short {fp + "0", false}, // one char long } for _, c := range cases { if got := validFingerprint(c.s); got != c.ok { t.Errorf("validFingerprint(%q) = %v, wanted %v", c.s, got, c.ok) } } } func TestPinnedClientTLSConfig_RejectsBadFingerprint(t *testing.T) { if _, err := pinnedClientTLSConfig("too short"); err == nil { t.Fatal("expected an error for a malformed fingerprint") } } func testCert(t *testing.T) (cert tls.Certificate, fingerprint string) { t.Helper() certDER, keyDER, err := generateHubCert() if err != nil { t.Fatal(err) } priv, err := x509.ParsePKCS8PrivateKey(keyDER) if err != nil { t.Fatal(err) } return tls.Certificate{Certificate: [][]byte{certDER}, PrivateKey: priv}, certFingerprint(certDER) } // listenTLS starts a TLS server on a random localhost port with the given // certificate and returns its address. Connections are drained (not // closed outright) so the TLS handshake — which crypto/tls only performs // lazily, on first Read/Write — actually gets a chance to complete. func listenTLS(t *testing.T, cert tls.Certificate) string { t.Helper() ln, err := tls.Listen("tcp", "127.0.0.1:0", hubServerTLSConfig(cert)) if err != nil { t.Fatal(err) } t.Cleanup(func() { ln.Close() }) go func() { for { c, err := ln.Accept() if err != nil { return } go func(c net.Conn) { defer c.Close() io.Copy(io.Discard, c) }(c) } }() return ln.Addr().String() } func TestPinnedClientTLSConfig_HandshakeEndToEnd(t *testing.T) { certA, fpA := testCert(t) _, fpB := testCert(t) // a different cert, never presented by the server addr := listenTLS(t, certA) // Correct fingerprint: handshake succeeds. cfg, err := pinnedClientTLSConfig(fpA) if err != nil { t.Fatal(err) } conn, err := tls.Dial("tcp", addr, cfg) if err != nil { t.Fatalf("expected the handshake to succeed with the right fingerprint: %v", err) } conn.Close() // Pinned to a fingerprint the server never presents: handshake must // fail, even though certB (fpB) is a perfectly valid certificate on // its own — it's just not the one at this address. cfgWrong, err := pinnedClientTLSConfig(fpB) if err != nil { t.Fatal(err) } if _, err := tls.Dial("tcp", addr, cfgWrong); err == nil { t.Fatal("expected the handshake to fail: server presented a cert that doesn't match the pinned fingerprint") } }