Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 25 additions & 12 deletions internal/atunnel/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -89,13 +89,9 @@ func NewServer(cfg Config) (*Server, error) {
if _, err := loadCredentialBundle(cfg.CredentialBundlePath); err != nil {
return nil, err
}
trustPEM, err := os.ReadFile(cfg.TrustBundlePath)
clientCAs, err := loadTrustBundle(cfg.TrustBundlePath)
if err != nil {
return nil, fmt.Errorf("atunnel: reading trust bundle: %w", err)
}
clientCAs := x509.NewCertPool()
if !clientCAs.AppendCertsFromPEM(trustPEM) {
return nil, fmt.Errorf("atunnel: trust bundle %q contains no certificates", cfg.TrustBundlePath)
return nil, err
}

proxy := httputil.NewSingleHostReverseProxy(cfg.Upstream)
Expand All @@ -115,13 +111,18 @@ func NewServer(cfg Config) (*Server, error) {
GetCertificate: func(*tls.ClientHelloInfo) (*tls.Certificate, error) {
return loadCredentialBundle(s.credentialBundlePath)
},
GetConfigForClient: func(*tls.ClientHelloInfo) (*tls.Config, error) {
clientCAs, err := loadTrustBundle(cfg.TrustBundlePath)
if err != nil {
return nil, err
}
config := s.tlsConfig.Clone()
config.GetConfigForClient = nil
config.ClientCAs = clientCAs
return config, nil
},
ClientAuth: tls.RequireAndVerifyClientCert,
// TODO(liorlieberman): reload the trust bundle per connection via
// GetConfigForClient, mirroring GetCertificate above. kubelet keeps the
// projected ClusterTrustBundle in sync with the signer, but this pool is
// frozen at process start, so after a CA rotation a long-lived worker
// rejects the router until its pod restarts.
ClientCAs: clientCAs,
ClientCAs: clientCAs,
VerifyConnection: func(cs tls.ConnectionState) error {
if len(cs.PeerCertificates) == 0 {
return fmt.Errorf("atunnel: client certificate is required")
Expand Down Expand Up @@ -149,6 +150,18 @@ func loadCredentialBundle(path string) (*tls.Certificate, error) {
return &cert, nil
}

func loadTrustBundle(path string) (*x509.CertPool, error) {
trustPEM, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("atunnel: reading trust bundle: %w", err)
}
clientCAs := x509.NewCertPool()
if !clientCAs.AppendCertsFromPEM(trustPEM) {
return nil, fmt.Errorf("atunnel: trust bundle %q contains no certificates", path)
}
return clientCAs, nil
}

// Serve serves HTTPS on lis until ctx is canceled or the server fails.
func (s *Server) Serve(ctx context.Context, lis net.Listener) error {
httpServer := &http.Server{
Expand Down
51 changes: 51 additions & 0 deletions internal/atunnel/server_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -227,6 +227,57 @@ func TestMutualTLSClientIdentity(t *testing.T) {
}
}

func TestMutualTLSReloadsClientTrustBundle(t *testing.T) {
dir := t.TempDir()
firstCA := newTestCA(t)
secondCA := newTestCA(t)
bundlePath := filepath.Join(dir, "server.pem")
trustPath := filepath.Join(dir, "trust.pem")
writeCredentialBundle(t, bundlePath, firstCA.issue(t, "", []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}))
if err := os.WriteFile(trustPath, firstCA.certPEM, 0o600); err != nil {
t.Fatal(err)
}
upstream, err := url.Parse("http://actor.internal:80")
if err != nil {
t.Fatal(err)
}
s, err := NewServer(Config{
CredentialBundlePath: bundlePath,
TrustBundlePath: trustPath,
AllowedClientID: "spiffe://cluster.local/ns/ate-system/sa/atenet-router",
Upstream: upstream,
})
if err != nil {
t.Fatal(err)
}

clientID := "spiffe://cluster.local/ns/ate-system/sa/atenet-router"
firstClient := firstCA.issue(t, clientID, []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth})
secondClient := secondCA.issue(t, clientID, []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth})
assertHandshake := func(name string, cert tls.Certificate, wantErr bool) {
t.Helper()
t.Run(name, func(t *testing.T) {
serverErr, clientErr := tlsHandshake(s.tlsConfig, &tls.Config{
MinVersion: tls.VersionTLS12,
InsecureSkipVerify: true, // Test only the server's client authentication here.
Certificates: []tls.Certificate{cert},
})
gotErr := serverErr != nil || clientErr != nil
if gotErr != wantErr {
t.Fatalf("server error = %v, client error = %v, want error %v", serverErr, clientErr, wantErr)
}
})
}

assertHandshake("original CA before rotation", firstClient, false)
assertHandshake("new CA before rotation", secondClient, true)
if err := os.WriteFile(trustPath, secondCA.certPEM, 0o600); err != nil {
t.Fatal(err)
}
assertHandshake("original CA after rotation", firstClient, true)
assertHandshake("new CA after rotation", secondClient, false)
}

func TestDeactivateCancelsInflightRequest(t *testing.T) {
upstream, err := url.Parse("http://actor.internal:80")
if err != nil {
Expand Down