package auth import ( "crypto/ed25519" "crypto/rand" "crypto/tls" "crypto/x509" "crypto/x509/pkix" "encoding/pem" "math/big" "net" "net/url" "os" "path/filepath" "testing" "time" ) type certFiles struct{ cert, key string } func writeCertificate(t *testing.T, dir, name string, template, parent *x509.Certificate, parentKey ed25519.PrivateKey) (certFiles, ed25519.PrivateKey) { t.Helper() public, private, err := ed25519.GenerateKey(rand.Reader) if err != nil { t.Fatal(err) } signer := parentKey if signer == nil { signer = private } der, err := x509.CreateCertificate(rand.Reader, template, parent, public, signer) if err != nil { t.Fatal(err) } certPath, keyPath := filepath.Join(dir, name+".crt"), filepath.Join(dir, name+".key") keyDER, err := x509.MarshalPKCS8PrivateKey(private) if err != nil { t.Fatal(err) } if err := os.WriteFile(certPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}), 0o600); err != nil { t.Fatal(err) } if err := os.WriteFile(keyPath, pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyDER}), 0o600); err != nil { t.Fatal(err) } return certFiles{cert: certPath, key: keyPath}, private } func makeTLSFixtures(t *testing.T) (caPath string, server, edge, wrongRole certFiles) { t.Helper() dir := t.TempDir() now := time.Now() caTemplate := &x509.Certificate{SerialNumber: big.NewInt(1), Subject: pkix.Name{CommonName: "test-ca"}, NotBefore: now.Add(-time.Hour), NotAfter: now.Add(time.Hour), IsCA: true, BasicConstraintsValid: true, KeyUsage: x509.KeyUsageCertSign} caFiles, caKey := writeCertificate(t, dir, "ca", caTemplate, caTemplate, nil) caPEM, err := os.ReadFile(caFiles.cert) if err != nil { t.Fatal(err) } caCertBlock, _ := pem.Decode(caPEM) caCert, err := x509.ParseCertificate(caCertBlock.Bytes) if err != nil { t.Fatal(err) } serverURI, _ := url.Parse("spiffe://iop/control-plane/cp-1") serverTemplate := &x509.Certificate{SerialNumber: big.NewInt(2), Subject: pkix.Name{CommonName: "control-plane"}, DNSNames: []string{"control-plane.internal"}, URIs: []*url.URL{serverURI}, NotBefore: now.Add(-time.Hour), NotAfter: now.Add(time.Hour), KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}} server, _ = writeCertificate(t, dir, "server", serverTemplate, caCert, caKey) edgeURI, _ := url.Parse("spiffe://iop/edge/edge-1") edgeTemplate := &x509.Certificate{SerialNumber: big.NewInt(3), Subject: pkix.Name{CommonName: "edge"}, URIs: []*url.URL{edgeURI}, NotBefore: now.Add(-time.Hour), NotAfter: now.Add(time.Hour), KeyUsage: x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}} edge, _ = writeCertificate(t, dir, "edge", edgeTemplate, caCert, caKey) nodeURI, _ := url.Parse("spiffe://iop/node/node-1") nodeTemplate := *edgeTemplate nodeTemplate.SerialNumber = big.NewInt(4) nodeTemplate.URIs = []*url.URL{nodeURI} wrongRole, _ = writeCertificate(t, dir, "node", &nodeTemplate, caCert, caKey) return caFiles.cert, server, edge, wrongRole } func handshake(serverConfig, clientConfig *tls.Config) (error, error) { serverConn, clientConn := net.Pipe() deadline := time.Now().Add(2 * time.Second) _ = serverConn.SetDeadline(deadline) _ = clientConn.SetDeadline(deadline) serverTLS, clientTLS := tls.Server(serverConn, serverConfig), tls.Client(clientConn, clientConfig) serverResult := make(chan error, 1) go func() { serverResult <- serverTLS.Handshake() }() clientErr := clientTLS.Handshake() serverErr := <-serverResult _ = serverConn.Close() _ = clientConn.Close() return serverErr, clientErr } func TestMutualTLS13PeerIdentityMatrix(t *testing.T) { ca, serverFiles, edgeFiles, wrongRoleFiles := makeTLSFixtures(t) serverConfig, err := LoadServerTLSWithIdentity(serverFiles.cert, serverFiles.key, ca, "edge", "") if err != nil { t.Fatal(err) } clientConfig, err := LoadClientTLSWithIdentity(edgeFiles.cert, edgeFiles.key, ca, "control-plane.internal", "control-plane", "cp-1") if err != nil { t.Fatal(err) } if serverErr, clientErr := handshake(serverConfig, clientConfig); serverErr != nil || clientErr != nil { t.Fatalf("valid mTLS failed: server=%v client=%v", serverErr, clientErr) } if clientConfig.MinVersion != tls.VersionTLS13 || serverConfig.MinVersion != tls.VersionTLS13 { t.Fatal("credential transport does not require TLS 1.3") } wrongRole, err := LoadClientTLSWithIdentity(wrongRoleFiles.cert, wrongRoleFiles.key, ca, "control-plane.internal", "control-plane", "cp-1") if err != nil { t.Fatal(err) } if serverErr, _ := handshake(serverConfig, wrongRole); serverErr == nil { t.Fatal("server accepted wrong peer role") } _, _, otherEdge, _ := makeTLSFixtures(t) wrongCA, err := LoadClientTLSWithIdentity(otherEdge.cert, otherEdge.key, ca, "control-plane.internal", "control-plane", "cp-1") if err != nil { t.Fatal(err) } if serverErr, _ := handshake(serverConfig, wrongCA); serverErr == nil { t.Fatal("server accepted client signed by wrong CA") } missingCertificate := clientConfig.Clone() missingCertificate.Certificates = nil if serverErr, _ := handshake(serverConfig, missingCertificate); serverErr == nil { t.Fatal("server accepted missing client certificate") } wrongName := clientConfig.Clone() wrongName.ServerName = "wrong.internal" if _, clientErr := handshake(serverConfig, wrongName); clientErr == nil { t.Fatal("client accepted wrong server name") } exactNameConfig, err := LoadServerTLSWithIdentity(serverFiles.cert, serverFiles.key, ca, "edge", "edge-1") if err != nil { t.Fatal(err) } if serverErr, clientErr := handshake(exactNameConfig, clientConfig); serverErr != nil || clientErr != nil { t.Fatalf("exact workload name failed: server=%v client=%v", serverErr, clientErr) } wrongWorkloadNameConfig, err := LoadServerTLSWithIdentity(serverFiles.cert, serverFiles.key, ca, "edge", "edge-other") if err != nil { t.Fatal(err) } if serverErr, _ := handshake(wrongWorkloadNameConfig, clientConfig); serverErr == nil { t.Fatal("server accepted same-role different workload name") } } func TestParseWorkloadIdentityRequiresOneCanonicalIdentity(t *testing.T) { parseURI := func(raw string) *url.URL { t.Helper() uri, err := url.Parse(raw) if err != nil { t.Fatal(err) } return uri } tests := []struct { name string uris []*url.URL want WorkloadIdentity wantErr bool }{ {name: "exact", uris: []*url.URL{parseURI("spiffe://example/ignored"), parseURI("spiffe://iop/edge/edge-1")}, want: WorkloadIdentity{Role: "edge", Name: "edge-1"}}, {name: "missing", uris: []*url.URL{parseURI("spiffe://example/edge/edge-1")}, wantErr: true}, {name: "malformed slash", uris: []*url.URL{parseURI("spiffe://iop/edge/team/edge-1")}, wantErr: true}, {name: "malformed query", uris: []*url.URL{parseURI("spiffe://iop/edge/edge-1?role=node")}, wantErr: true}, {name: "ambiguous", uris: []*url.URL{parseURI("spiffe://iop/edge/edge-1"), parseURI("spiffe://iop/edge/edge-2")}, wantErr: true}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { got, err := ParseWorkloadIdentity(&x509.Certificate{URIs: test.uris}) if test.wantErr { if err == nil { t.Fatalf("expected rejection, got %+v", got) } return } if err != nil { t.Fatal(err) } if got != test.want { t.Fatalf("identity: got %+v want %+v", got, test.want) } }) } }