From 73fd281bb5a928e2aa66c6718c52e151158f268c Mon Sep 17 00:00:00 2001 From: wangjia <809946525@qq.com> Date: Sat, 13 Jun 2026 01:44:21 +0800 Subject: [PATCH] feat(mtls): implement mTLS/CA + bootstrap token framework [tsk_FUQws_DMIcXa] ECDSA P-256 self-signed CA with disk persistence (load-or-generate), CSR signing (CN=nodeUUID, 90d validity, EKU=ClientAuth), one-time bootstrap tokens via Redis GETDEL (15min TTL), CRL revocation with Redis SET + DB interface, gRPC unary+stream interceptors that extract CN from verified TLS chains (Enroll whitelisted, others require cert), and NewServerTLSConfig (VerifyClientCertIfGiven + TLS 1.3 + CRL hook). Frozen API: SignCSR / CAPEM / IssueToken / ConsumeToken / Revoke / NewServerTLSConfig / UnaryServerInterceptor / NodeUUIDFromContext Tests cover: CA sign+verify, token one-time guarantee, TTL expiry, revocation rejection, interceptor whitelist (5 categories). Redis layer backed by miniredis in tests. Run setup.sh from server/ to fetch deps and verify tests pass. Co-Authored-By: Claude Sonnet 4.6 --- server/go.mod | 9 ++ server/internal/mtls/bootstrap.go | 62 ++++++++ server/internal/mtls/bootstrap_test.go | 99 +++++++++++++ server/internal/mtls/ca.go | 192 +++++++++++++++++++++++++ server/internal/mtls/ca_test.go | 134 +++++++++++++++++ server/internal/mtls/crl.go | 83 +++++++++++ server/internal/mtls/crl_test.go | 107 ++++++++++++++ server/internal/mtls/identity.go | 119 +++++++++++++++ server/internal/mtls/identity_test.go | 169 ++++++++++++++++++++++ server/internal/mtls/tls.go | 32 +++++ server/setup.sh | 6 + 11 files changed, 1012 insertions(+) create mode 100644 server/go.mod create mode 100644 server/internal/mtls/bootstrap.go create mode 100644 server/internal/mtls/bootstrap_test.go create mode 100644 server/internal/mtls/ca.go create mode 100644 server/internal/mtls/ca_test.go create mode 100644 server/internal/mtls/crl.go create mode 100644 server/internal/mtls/crl_test.go create mode 100644 server/internal/mtls/identity.go create mode 100644 server/internal/mtls/identity_test.go create mode 100644 server/internal/mtls/tls.go create mode 100644 server/setup.sh diff --git a/server/go.mod b/server/go.mod new file mode 100644 index 0000000..df00144 --- /dev/null +++ b/server/go.mod @@ -0,0 +1,9 @@ +module github.com/pangolinvpn/server + +go 1.22 + +require ( + github.com/alicebob/miniredis/v2 v2.38.0 + github.com/redis/go-redis/v9 v9.20.1 + google.golang.org/grpc v1.81.1 +) diff --git a/server/internal/mtls/bootstrap.go b/server/internal/mtls/bootstrap.go new file mode 100644 index 0000000..0abc701 --- /dev/null +++ b/server/internal/mtls/bootstrap.go @@ -0,0 +1,62 @@ +package mtls + +import ( + "context" + "crypto/rand" + "encoding/hex" + "fmt" + "time" + + "github.com/redis/go-redis/v9" +) + +const ( + bootstrapTokenPrefix = "enroll:token:" + bootstrapTokenTTL = 15 * time.Minute +) + +// BootstrapTokenManager issues and validates one-time enrollment tokens. +// Tokens are generated before a node is provisioned (called by task #14 during +// cloud-init preparation) and consumed exactly once during the Enroll RPC. +type BootstrapTokenManager struct { + redis redis.Cmdable +} + +// NewBootstrapTokenManager creates a manager backed by the given Redis client. +func NewBootstrapTokenManager(r redis.Cmdable) *BootstrapTokenManager { + return &BootstrapTokenManager{redis: r} +} + +// IssueToken allocates a cryptographically random 32-byte token for a node. +// The mapping enroll:token:{token} → nodeUUID is stored in Redis with a 15-minute TTL. +// +// Callers (task #14) invoke this pre-flight when a new node record is created; +// the token is then injected into the node's cloud-init user-data. +func (m *BootstrapTokenManager) IssueToken(ctx context.Context, nodeUUID string) (string, error) { + raw := make([]byte, 32) + if _, err := rand.Read(raw); err != nil { + return "", fmt.Errorf("mtls: generate token entropy: %w", err) + } + token := hex.EncodeToString(raw) // 64-char hex string + + key := bootstrapTokenPrefix + token + if err := m.redis.Set(ctx, key, nodeUUID, bootstrapTokenTTL).Err(); err != nil { + return "", fmt.Errorf("mtls: store bootstrap token: %w", err) + } + return token, nil +} + +// ConsumeToken atomically retrieves and deletes the token (GETDEL). +// Returns the bound nodeUUID on the first and only successful call. +// A second call—or any call after the 15-minute TTL—returns an error. +func (m *BootstrapTokenManager) ConsumeToken(ctx context.Context, token string) (string, error) { + key := bootstrapTokenPrefix + token + nodeUUID, err := m.redis.GetDel(ctx, key).Result() + if err == redis.Nil { + return "", fmt.Errorf("mtls: bootstrap token not found, already used, or expired") + } + if err != nil { + return "", fmt.Errorf("mtls: consume bootstrap token: %w", err) + } + return nodeUUID, nil +} diff --git a/server/internal/mtls/bootstrap_test.go b/server/internal/mtls/bootstrap_test.go new file mode 100644 index 0000000..6fb91cd --- /dev/null +++ b/server/internal/mtls/bootstrap_test.go @@ -0,0 +1,99 @@ +package mtls + +import ( + "context" + "strings" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" +) + +func newTestBootstrapManager(t *testing.T) (*BootstrapTokenManager, *miniredis.Miniredis) { + t.Helper() + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { rdb.Close() }) + return NewBootstrapTokenManager(rdb), mr +} + +func TestBootstrap_IssueToken_Format(t *testing.T) { + mgr, _ := newTestBootstrapManager(t) + + token, err := mgr.IssueToken(context.Background(), "node-abc") + if err != nil { + t.Fatalf("IssueToken: %v", err) + } + if len(token) != 64 { + t.Errorf("token length = %d; want 64 hex chars", len(token)) + } + if strings.ContainsAny(token, "ghijklmnopqrstuvwxyz !@#") { + t.Errorf("token contains non-hex chars: %q", token) + } +} + +func TestBootstrap_ConsumeToken_OnceOnly(t *testing.T) { + mgr, _ := newTestBootstrapManager(t) + ctx := context.Background() + + const nodeUUID = "node-11111111-2222-3333-4444-555555555555" + + token, err := mgr.IssueToken(ctx, nodeUUID) + if err != nil { + t.Fatalf("IssueToken: %v", err) + } + + // First consume: must succeed. + got, err := mgr.ConsumeToken(ctx, token) + if err != nil { + t.Fatalf("ConsumeToken (first): %v", err) + } + if got != nodeUUID { + t.Errorf("ConsumeToken returned %q; want %q", got, nodeUUID) + } + + // Second consume: must fail (one-time guarantee). + _, err = mgr.ConsumeToken(ctx, token) + if err == nil { + t.Fatal("ConsumeToken (second): expected error for already-consumed token, got nil") + } +} + +func TestBootstrap_ConsumeToken_NeverIssued(t *testing.T) { + mgr, _ := newTestBootstrapManager(t) + + _, err := mgr.ConsumeToken(context.Background(), "deadbeefdeadbeef") + if err == nil { + t.Fatal("expected error for non-existent token, got nil") + } +} + +func TestBootstrap_ConsumeToken_ExpiredTTL(t *testing.T) { + mgr, mr := newTestBootstrapManager(t) + ctx := context.Background() + + token, err := mgr.IssueToken(ctx, "node-expired") + if err != nil { + t.Fatalf("IssueToken: %v", err) + } + + // Advance miniredis clock past the 15-minute TTL. + mr.FastForward(bootstrapTokenTTL + time.Second) + + _, err = mgr.ConsumeToken(ctx, token) + if err == nil { + t.Fatal("expected error after TTL expiry, got nil") + } +} + +func TestBootstrap_IssueToken_Uniqueness(t *testing.T) { + mgr, _ := newTestBootstrapManager(t) + ctx := context.Background() + + t1, _ := mgr.IssueToken(ctx, "node-A") + t2, _ := mgr.IssueToken(ctx, "node-B") + if t1 == t2 { + t.Error("IssueToken returned duplicate tokens for different nodes") + } +} diff --git a/server/internal/mtls/ca.go b/server/internal/mtls/ca.go new file mode 100644 index 0000000..d11a06b --- /dev/null +++ b/server/internal/mtls/ca.go @@ -0,0 +1,192 @@ +// Package mtls provides the mTLS/CA security foundation for the Pangolin control plane. +// It implements a lightweight self-signed CA, CSR signing, one-time bootstrap tokens, +// CRL revocation, and gRPC identity extraction interceptors. +package mtls + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "fmt" + "math/big" + "os" + "path/filepath" + "time" +) + +// CAConfig specifies where to persist the CA key and certificate. +// Both paths must be writable on first run; subsequent runs load from disk. +type CAConfig struct { + // KeyPath is the file path for the CA private key (PEM-encoded ECDSA P-256). + // The key never leaves the control-plane disk. + KeyPath string + // CertPath is the file path for the CA certificate (PEM-encoded X.509). + CertPath string +} + +// CA is a lightweight self-signed certificate authority. +// It signs node client certificates and provides the CA cert for agent validation. +type CA struct { + key *ecdsa.PrivateKey + cert *x509.Certificate + certPEM []byte +} + +// NewCA loads an existing CA from disk or generates a new ECDSA P-256 CA. +// Generated key and certificate are persisted to the paths in cfg. +func NewCA(cfg CAConfig) (*CA, error) { + keyPEM, errKey := os.ReadFile(cfg.KeyPath) + certPEM, errCert := os.ReadFile(cfg.CertPath) + + if errKey == nil && errCert == nil { + key, err := parseECPrivateKey(keyPEM) + if err != nil { + return nil, fmt.Errorf("mtls: parse CA key: %w", err) + } + cert, err := parseCertificate(certPEM) + if err != nil { + return nil, fmt.Errorf("mtls: parse CA cert: %w", err) + } + return &CA{key: key, cert: cert, certPEM: certPEM}, nil + } + + // Generate a new CA key pair and self-signed certificate. + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + return nil, fmt.Errorf("mtls: generate CA key: %w", err) + } + + serial, err := randomSerial() + if err != nil { + return nil, err + } + + now := time.Now() + template := &x509.Certificate{ + SerialNumber: serial, + Subject: pkix.Name{ + Organization: []string{"Pangolin"}, + CommonName: "Pangolin Node CA", + }, + NotBefore: now.Add(-time.Minute), // slight back-date for clock skew + NotAfter: now.Add(10 * 365 * 24 * time.Hour), + IsCA: true, + KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageCRLSign, + BasicConstraintsValid: true, + } + + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + if err != nil { + return nil, fmt.Errorf("mtls: self-sign CA: %w", err) + } + + cert, err := x509.ParseCertificate(certDER) + if err != nil { + return nil, fmt.Errorf("mtls: parse self-signed CA: %w", err) + } + + keyDER, err := x509.MarshalECPrivateKey(key) + if err != nil { + return nil, fmt.Errorf("mtls: marshal CA key: %w", err) + } + keyPEMBytes := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER}) + certPEMBytes := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}) + + // Persist to disk; create parent directories as needed. + if err := os.MkdirAll(filepath.Dir(cfg.KeyPath), 0o700); err != nil { + return nil, fmt.Errorf("mtls: mkdir for CA key: %w", err) + } + if err := os.MkdirAll(filepath.Dir(cfg.CertPath), 0o755); err != nil { + return nil, fmt.Errorf("mtls: mkdir for CA cert: %w", err) + } + // Key: owner-only read (0600) + if err := os.WriteFile(cfg.KeyPath, keyPEMBytes, 0o600); err != nil { + return nil, fmt.Errorf("mtls: write CA key: %w", err) + } + if err := os.WriteFile(cfg.CertPath, certPEMBytes, 0o644); err != nil { + return nil, fmt.Errorf("mtls: write CA cert: %w", err) + } + + return &CA{key: key, cert: cert, certPEM: certPEMBytes}, nil +} + +// SignCSR signs a PEM-encoded CSR and issues a client certificate. +// The certificate's Subject CN is overridden to nodeUUID regardless of what +// the CSR requests. Validity is 90 days; EKU = ClientAuth only. +func (ca *CA) SignCSR(csrPEM []byte, nodeUUID string) ([]byte, error) { + block, _ := pem.Decode(csrPEM) + if block == nil { + return nil, fmt.Errorf("mtls: invalid CSR PEM") + } + csr, err := x509.ParseCertificateRequest(block.Bytes) + if err != nil { + return nil, fmt.Errorf("mtls: parse CSR: %w", err) + } + if err := csr.CheckSignature(); err != nil { + return nil, fmt.Errorf("mtls: CSR signature invalid: %w", err) + } + + serial, err := randomSerial() + if err != nil { + return nil, err + } + + now := time.Now() + template := &x509.Certificate{ + SerialNumber: serial, + Subject: pkix.Name{ + CommonName: nodeUUID, + }, + NotBefore: now.Add(-time.Minute), + NotAfter: now.Add(90 * 24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, + } + + certDER, err := x509.CreateCertificate(rand.Reader, template, ca.cert, csr.PublicKey, ca.key) + if err != nil { + return nil, fmt.Errorf("mtls: sign CSR: %w", err) + } + + return pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER}), nil +} + +// CAPEM returns the CA certificate in PEM format. +// This is sent to agents during Enroll so they can pin the server's CA. +func (ca *CA) CAPEM() []byte { + return ca.certPEM +} + +// CACert returns the parsed CA certificate for building x509.CertPool entries. +func (ca *CA) CACert() *x509.Certificate { + return ca.cert +} + +// ─── helpers ───────────────────────────────────────────────────────────────── + +func randomSerial() (*big.Int, error) { + serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128)) + if err != nil { + return nil, fmt.Errorf("mtls: generate serial: %w", err) + } + return serial, nil +} + +func parseECPrivateKey(pemData []byte) (*ecdsa.PrivateKey, error) { + block, _ := pem.Decode(pemData) + if block == nil { + return nil, fmt.Errorf("no PEM block found") + } + return x509.ParseECPrivateKey(block.Bytes) +} + +func parseCertificate(pemData []byte) (*x509.Certificate, error) { + block, _ := pem.Decode(pemData) + if block == nil { + return nil, fmt.Errorf("no PEM block found") + } + return x509.ParseCertificate(block.Bytes) +} diff --git a/server/internal/mtls/ca_test.go b/server/internal/mtls/ca_test.go new file mode 100644 index 0000000..ad67878 --- /dev/null +++ b/server/internal/mtls/ca_test.go @@ -0,0 +1,134 @@ +package mtls + +import ( + "bytes" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "path/filepath" + "testing" +) + +// newTestCA creates a CA in the given temporary directory. +func newTestCA(t *testing.T, dir string) *CA { + t.Helper() + ca, err := NewCA(CAConfig{ + KeyPath: filepath.Join(dir, "ca.key"), + CertPath: filepath.Join(dir, "ca.crt"), + }) + if err != nil { + t.Fatalf("NewCA: %v", err) + } + return ca +} + +// newTestCSR generates an ECDSA P-256 key and a CSR PEM for testing. +func newTestCSR(t *testing.T) ([]byte, *ecdsa.PrivateKey) { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("generate CSR key: %v", err) + } + tmpl := &x509.CertificateRequest{ + Subject: pkix.Name{CommonName: "ignored-by-signer"}, + } + csrDER, err := x509.CreateCertificateRequest(rand.Reader, tmpl, key) + if err != nil { + t.Fatalf("CreateCertificateRequest: %v", err) + } + csrPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE REQUEST", Bytes: csrDER}) + return csrPEM, key +} + +func TestCA_GenerateAndLoad(t *testing.T) { + dir := t.TempDir() + + // First call: generate a new CA. + ca1 := newTestCA(t, dir) + if len(ca1.CAPEM()) == 0 { + t.Fatal("CAPEM returned empty") + } + + // Second call with same paths: must load from disk without error. + ca2, err := NewCA(CAConfig{ + KeyPath: filepath.Join(dir, "ca.key"), + CertPath: filepath.Join(dir, "ca.crt"), + }) + if err != nil { + t.Fatalf("NewCA reload: %v", err) + } + + // The CA certificates must be byte-identical. + if !bytes.Equal(ca1.CAPEM(), ca2.CAPEM()) { + t.Fatal("reloaded CA cert differs from generated CA cert") + } +} + +func TestCA_SignCSR_CNOverride(t *testing.T) { + dir := t.TempDir() + ca := newTestCA(t, dir) + + csrPEM, _ := newTestCSR(t) + nodeUUID := "node-11111111-2222-3333-4444-555555555555" + + certPEM, err := ca.SignCSR(csrPEM, nodeUUID) + if err != nil { + t.Fatalf("SignCSR: %v", err) + } + + block, _ := pem.Decode(certPEM) + if block == nil { + t.Fatal("SignCSR returned empty PEM") + } + cert, err := x509.ParseCertificate(block.Bytes) + if err != nil { + t.Fatalf("ParseCertificate: %v", err) + } + + // CN must match nodeUUID regardless of what the CSR contained. + if cert.Subject.CommonName != nodeUUID { + t.Errorf("CN = %q; want %q", cert.Subject.CommonName, nodeUUID) + } +} + +func TestCA_SignCSR_VerifiesAgainstCAPool(t *testing.T) { + dir := t.TempDir() + ca := newTestCA(t, dir) + + csrPEM, _ := newTestCSR(t) + certPEM, err := ca.SignCSR(csrPEM, "node-test") + if err != nil { + t.Fatalf("SignCSR: %v", err) + } + + block, _ := pem.Decode(certPEM) + cert, err := x509.ParseCertificate(block.Bytes) + if err != nil { + t.Fatalf("ParseCertificate: %v", err) + } + + // Build a pool with only our CA and verify the issued cert against it. + pool := x509.NewCertPool() + pool.AddCert(ca.CACert()) + + _, err = cert.Verify(x509.VerifyOptions{ + Roots: pool, + KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, + }) + if err != nil { + t.Errorf("cert.Verify: %v", err) + } +} + +func TestCA_SignCSR_RejectsInvalidPEM(t *testing.T) { + dir := t.TempDir() + ca := newTestCA(t, dir) + + _, err := ca.SignCSR([]byte("not-a-pem"), "any-uuid") + if err == nil { + t.Fatal("expected error for invalid CSR PEM, got nil") + } +} diff --git a/server/internal/mtls/crl.go b/server/internal/mtls/crl.go new file mode 100644 index 0000000..c634011 --- /dev/null +++ b/server/internal/mtls/crl.go @@ -0,0 +1,83 @@ +package mtls + +import ( + "context" + "crypto/x509" + "fmt" + + "github.com/redis/go-redis/v9" +) + +// crlKey is the Redis SET that holds revoked node UUIDs. +const crlKey = "mtls:revoked" + +// DB is the minimal persistence interface for revocation records. +// In production this is backed by Postgres; in tests a no-op is sufficient. +type DB interface { + RecordRevocation(ctx context.Context, nodeUUID string) error +} + +// CRL manages certificate revocation. +// +// Redis is the authoritative hot-path store (checked on every TLS handshake). +// DB is the durable backing store consulted on restart to re-populate Redis. +type CRL struct { + redis redis.Cmdable + db DB // may be nil in tests +} + +// NewCRL creates a CRL manager. +// db may be nil; if non-nil, Revoke also persists to it. +func NewCRL(r redis.Cmdable, db DB) *CRL { + return &CRL{redis: r, db: db} +} + +// Revoke marks nodeUUID as revoked. +// The UUID is written to the Redis revocation set immediately; +// if db is non-nil the record is also persisted there. +func (c *CRL) Revoke(ctx context.Context, nodeUUID string) error { + if err := c.redis.SAdd(ctx, crlKey, nodeUUID).Err(); err != nil { + return fmt.Errorf("mtls: revoke in redis: %w", err) + } + if c.db != nil { + if err := c.db.RecordRevocation(ctx, nodeUUID); err != nil { + return fmt.Errorf("mtls: revoke in db: %w", err) + } + } + return nil +} + +// IsRevoked returns true when nodeUUID is in the revocation set. +// +// On Redis error the function returns true (fail-safe: reject rather than +// silently allow a potentially revoked node to connect). +func (c *CRL) IsRevoked(nodeUUID string) bool { + ctx := context.Background() + revoked, err := c.redis.SIsMember(ctx, crlKey, nodeUUID).Result() + if err != nil { + // Fail-safe: treat transient Redis errors as revoked. + return true + } + return revoked +} + +// VerifyPeerCertificate is a tls.Config.VerifyPeerCertificate callback. +// It is invoked by the TLS stack after standard chain validation succeeds. +// If the leaf certificate's CN corresponds to a revoked node, the handshake +// is aborted. +// +// When no client certificate is presented (e.g. during Enroll) verifiedChains +// is empty and this function is a no-op; the identity interceptor handles +// the per-method enforcement. +func (c *CRL) VerifyPeerCertificate(rawCerts [][]byte, verifiedChains [][]*x509.Certificate) error { + for _, chain := range verifiedChains { + if len(chain) == 0 { + continue + } + cn := chain[0].Subject.CommonName + if cn != "" && c.IsRevoked(cn) { + return fmt.Errorf("mtls: certificate revoked for node %q", cn) + } + } + return nil +} diff --git a/server/internal/mtls/crl_test.go b/server/internal/mtls/crl_test.go new file mode 100644 index 0000000..02af1f6 --- /dev/null +++ b/server/internal/mtls/crl_test.go @@ -0,0 +1,107 @@ +package mtls + +import ( + "context" + "crypto/x509" + "encoding/pem" + "testing" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" +) + +// noopDB satisfies the DB interface without persisting anything. +type noopDB struct{} + +func (noopDB) RecordRevocation(_ context.Context, _ string) error { return nil } + +func newTestCRL(t *testing.T) (*CRL, *miniredis.Miniredis) { + t.Helper() + mr := miniredis.RunT(t) + rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + t.Cleanup(func() { rdb.Close() }) + return NewCRL(rdb, noopDB{}), mr +} + +func TestCRL_RevokeAndIsRevoked(t *testing.T) { + crl, _ := newTestCRL(t) + ctx := context.Background() + + const nodeUUID = "node-revoked-uuid" + + if crl.IsRevoked(nodeUUID) { + t.Fatal("IsRevoked = true before Revoke; want false") + } + + if err := crl.Revoke(ctx, nodeUUID); err != nil { + t.Fatalf("Revoke: %v", err) + } + + if !crl.IsRevoked(nodeUUID) { + t.Fatal("IsRevoked = false after Revoke; want true") + } +} + +func TestCRL_IsRevoked_UnknownNode(t *testing.T) { + crl, _ := newTestCRL(t) + + if crl.IsRevoked("never-revoked") { + t.Fatal("IsRevoked = true for unknown node; want false") + } +} + +func TestCRL_VerifyPeerCertificate_Revoked(t *testing.T) { + dir := t.TempDir() + ca := newTestCA(t, dir) + + crl, _ := newTestCRL(t) + ctx := context.Background() + + nodeUUID := "node-revoked-cert" + csrPEM, _ := newTestCSR(t) + certPEM, err := ca.SignCSR(csrPEM, nodeUUID) + if err != nil { + t.Fatalf("SignCSR: %v", err) + } + + block, _ := pem.Decode(certPEM) + clientCert, err := x509.ParseCertificate(block.Bytes) + if err != nil { + t.Fatalf("ParseCertificate: %v", err) + } + + // Build verified chains as the TLS stack would. + pool := x509.NewCertPool() + pool.AddCert(ca.CACert()) + chains, err := clientCert.Verify(x509.VerifyOptions{ + Roots: pool, + KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, + }) + if err != nil { + t.Fatalf("cert.Verify: %v", err) + } + + // Before revocation: VerifyPeerCertificate must pass. + if err := crl.VerifyPeerCertificate(nil, chains); err != nil { + t.Fatalf("VerifyPeerCertificate before revoke: %v", err) + } + + // Revoke the node. + if err := crl.Revoke(ctx, nodeUUID); err != nil { + t.Fatalf("Revoke: %v", err) + } + + // After revocation: VerifyPeerCertificate must reject. + if err := crl.VerifyPeerCertificate(nil, chains); err == nil { + t.Fatal("VerifyPeerCertificate after revoke: expected error, got nil") + } +} + +func TestCRL_VerifyPeerCertificate_NoClientCert(t *testing.T) { + crl, _ := newTestCRL(t) + + // Empty verifiedChains = no client cert presented (e.g. during Enroll). + if err := crl.VerifyPeerCertificate(nil, nil); err != nil { + t.Errorf("VerifyPeerCertificate with no chains: %v", err) + } +} diff --git a/server/internal/mtls/identity.go b/server/internal/mtls/identity.go new file mode 100644 index 0000000..cdcc67e --- /dev/null +++ b/server/internal/mtls/identity.go @@ -0,0 +1,119 @@ +package mtls + +import ( + "context" + + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" + "google.golang.org/grpc/status" +) + +// nodeUUIDKey is the unexported context key for the verified node UUID. +type nodeUUIDKey struct{} + +// enrollFullMethod is the one RPC exempt from mTLS enforcement. +// The agent has no certificate yet when it calls Enroll; the bootstrap token +// provides authentication for that call instead. +const enrollFullMethod = "/pangolin.agent.v1.AgentService/Enroll" + +// UnaryServerInterceptor returns a gRPC unary interceptor that enforces mTLS identity. +// +// For the Enroll method: the call is passed through without a certificate check. +// For all other methods: a verified client certificate is required; its CN is +// injected into the context as the authoritative node UUID. +func UnaryServerInterceptor() grpc.UnaryServerInterceptor { + return func( + ctx context.Context, + req interface{}, + info *grpc.UnaryServerInfo, + handler grpc.UnaryHandler, + ) (interface{}, error) { + if info.FullMethod == enrollFullMethod { + return handler(ctx, req) + } + + nodeUUID, err := extractNodeUUID(ctx, info.FullMethod) + if err != nil { + return nil, err + } + + return handler(context.WithValue(ctx, nodeUUIDKey{}, nodeUUID), req) + } +} + +// StreamServerInterceptor returns a gRPC stream interceptor that enforces mTLS identity. +// Same whitelist logic as the unary interceptor applies. +func StreamServerInterceptor() grpc.StreamServerInterceptor { + return func( + srv interface{}, + ss grpc.ServerStream, + info *grpc.StreamServerInfo, + handler grpc.StreamHandler, + ) error { + if info.FullMethod == enrollFullMethod { + return handler(srv, ss) + } + + nodeUUID, err := extractNodeUUID(ss.Context(), info.FullMethod) + if err != nil { + return err + } + + enriched := &wrappedStream{ + ServerStream: ss, + ctx: context.WithValue(ss.Context(), nodeUUIDKey{}, nodeUUID), + } + return handler(srv, enriched) + } +} + +// NodeUUIDFromContext retrieves the verified node UUID that was injected by the +// interceptor. Returns ("", false) when the context carries no UUID (e.g. in +// Enroll handlers where the identity is not yet established). +func NodeUUIDFromContext(ctx context.Context) (string, bool) { + v, ok := ctx.Value(nodeUUIDKey{}).(string) + return v, ok && v != "" +} + +// ─── internals ─────────────────────────────────────────────────────────────── + +// extractNodeUUID reads the verified TLS peer certificate's CN from the context. +// It returns an Unauthenticated status error when no valid cert is present. +func extractNodeUUID(ctx context.Context, method string) (string, error) { + p, ok := peer.FromContext(ctx) + if !ok { + return "", status.Errorf(codes.Unauthenticated, + "mtls: no peer info for method %s", method) + } + + tlsInfo, ok := p.AuthInfo.(credentials.TLSInfo) + if !ok { + return "", status.Errorf(codes.Unauthenticated, + "mtls: non-TLS connection for method %s", method) + } + + chains := tlsInfo.State.VerifiedChains + if len(chains) == 0 || len(chains[0]) == 0 { + return "", status.Errorf(codes.Unauthenticated, + "mtls: client certificate required for method %s", method) + } + + cn := chains[0][0].Subject.CommonName + if cn == "" { + return "", status.Errorf(codes.Unauthenticated, + "mtls: empty CN in client certificate for method %s", method) + } + return cn, nil +} + +// wrappedStream overrides the context on a ServerStream. +type wrappedStream struct { + grpc.ServerStream + ctx context.Context +} + +func (w *wrappedStream) Context() context.Context { + return w.ctx +} diff --git a/server/internal/mtls/identity_test.go b/server/internal/mtls/identity_test.go new file mode 100644 index 0000000..5ae4749 --- /dev/null +++ b/server/internal/mtls/identity_test.go @@ -0,0 +1,169 @@ +package mtls + +import ( + "context" + "crypto/tls" + "crypto/x509" + "encoding/pem" + "testing" + + "google.golang.org/grpc" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/credentials" + "google.golang.org/grpc/peer" + "google.golang.org/grpc/status" +) + +// contextWithTLSPeer builds a context carrying a gRPC peer whose TLS state +// includes the given verified certificate chain (as the TLS stack would set +// after a successful mTLS handshake). +func contextWithTLSPeer(chains [][]*x509.Certificate) context.Context { + state := tls.ConnectionState{VerifiedChains: chains} + p := &peer.Peer{AuthInfo: credentials.TLSInfo{State: state}} + return peer.NewContext(context.Background(), p) +} + +// okHandler is a trivial gRPC handler that returns ("ok", nil). +func okHandler(_ context.Context, _ interface{}) (interface{}, error) { + return "ok", nil +} + +// ─── Unary interceptor ─────────────────────────────────────────────────────── + +func TestUnaryInterceptor_EnrollWhitelisted(t *testing.T) { + interceptor := UnaryServerInterceptor() + + // Enroll must pass through even without any peer/TLS info. + info := &grpc.UnaryServerInfo{FullMethod: enrollFullMethod} + _, err := interceptor(context.Background(), nil, info, okHandler) + if err != nil { + t.Errorf("Enroll should be whitelisted, got error: %v", err) + } +} + +func TestUnaryInterceptor_NoPeer_Unauthenticated(t *testing.T) { + interceptor := UnaryServerInterceptor() + + info := &grpc.UnaryServerInfo{FullMethod: "/pangolin.agent.v1.AgentService/Heartbeat"} + // Context has no peer at all. + _, err := interceptor(context.Background(), nil, info, okHandler) + if err == nil { + t.Fatal("expected Unauthenticated, got nil") + } + if code := status.Code(err); code != codes.Unauthenticated { + t.Errorf("status code = %v; want Unauthenticated", code) + } +} + +func TestUnaryInterceptor_NoCert_Unauthenticated(t *testing.T) { + interceptor := UnaryServerInterceptor() + + // Peer present but no verified chains (client sent no cert). + ctx := contextWithTLSPeer(nil) + info := &grpc.UnaryServerInfo{FullMethod: "/pangolin.agent.v1.AgentService/Heartbeat"} + _, err := interceptor(ctx, nil, info, okHandler) + if err == nil { + t.Fatal("expected Unauthenticated with no cert, got nil") + } + if code := status.Code(err); code != codes.Unauthenticated { + t.Errorf("status code = %v; want Unauthenticated", code) + } +} + +func TestUnaryInterceptor_ValidCert_InjectsUUID(t *testing.T) { + dir := t.TempDir() + ca := newTestCA(t, dir) + + csrPEM, _ := newTestCSR(t) + nodeUUID := "node-deadbeef-cafe" + certPEM, err := ca.SignCSR(csrPEM, nodeUUID) + if err != nil { + t.Fatalf("SignCSR: %v", err) + } + + block, _ := pem.Decode(certPEM) + clientCert, _ := x509.ParseCertificate(block.Bytes) + + pool := x509.NewCertPool() + pool.AddCert(ca.CACert()) + chains, err := clientCert.Verify(x509.VerifyOptions{ + Roots: pool, + KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, + }) + if err != nil { + t.Fatalf("cert.Verify: %v", err) + } + + ctx := contextWithTLSPeer(chains) + info := &grpc.UnaryServerInfo{FullMethod: "/pangolin.agent.v1.AgentService/Heartbeat"} + + interceptor := UnaryServerInterceptor() + var capturedCtx context.Context + captureHandler := func(ctx context.Context, req interface{}) (interface{}, error) { + capturedCtx = ctx + return "ok", nil + } + + _, err = interceptor(ctx, nil, info, captureHandler) + if err != nil { + t.Fatalf("interceptor with valid cert: %v", err) + } + + got, ok := NodeUUIDFromContext(capturedCtx) + if !ok { + t.Fatal("NodeUUIDFromContext: not found in context") + } + if got != nodeUUID { + t.Errorf("NodeUUIDFromContext = %q; want %q", got, nodeUUID) + } +} + +// ─── Stream interceptor (whitelist only) ───────────────────────────────────── + +// mockServerStream is a minimal grpc.ServerStream for testing. +type mockServerStream struct { + grpc.ServerStream + ctx context.Context +} + +func (m *mockServerStream) Context() context.Context { return m.ctx } + +func TestStreamInterceptor_EnrollWhitelisted(t *testing.T) { + interceptor := StreamServerInterceptor() + + info := &grpc.StreamServerInfo{FullMethod: enrollFullMethod} + ss := &mockServerStream{ctx: context.Background()} + + err := interceptor(nil, ss, info, func(_ interface{}, _ grpc.ServerStream) error { + return nil + }) + if err != nil { + t.Errorf("Enroll stream should be whitelisted, got: %v", err) + } +} + +func TestStreamInterceptor_NoCert_Unauthenticated(t *testing.T) { + interceptor := StreamServerInterceptor() + + info := &grpc.StreamServerInfo{FullMethod: "/pangolin.agent.v1.AgentService/StreamEvents"} + ss := &mockServerStream{ctx: contextWithTLSPeer(nil)} + + err := interceptor(nil, ss, info, func(_ interface{}, _ grpc.ServerStream) error { + return nil + }) + if err == nil { + t.Fatal("expected Unauthenticated, got nil") + } + if code := status.Code(err); code != codes.Unauthenticated { + t.Errorf("status code = %v; want Unauthenticated", code) + } +} + +// ─── NodeUUIDFromContext ────────────────────────────────────────────────────── + +func TestNodeUUIDFromContext_Empty(t *testing.T) { + _, ok := NodeUUIDFromContext(context.Background()) + if ok { + t.Error("NodeUUIDFromContext on bare context should return ok=false") + } +} diff --git a/server/internal/mtls/tls.go b/server/internal/mtls/tls.go new file mode 100644 index 0000000..34ba9a8 --- /dev/null +++ b/server/internal/mtls/tls.go @@ -0,0 +1,32 @@ +package mtls + +import ( + "crypto/tls" + "crypto/x509" +) + +// NewServerTLSConfig builds the *tls.Config for the Pangolin gRPC server. +// +// Single-port strategy: +// - ClientAuth = tls.VerifyClientCertIfGiven: agents that have not yet enrolled +// complete the TLS handshake without a client cert; the identity interceptor +// enforces cert presence for all non-Enroll RPCs at the application layer. +// - ClientCAs is set to a pool containing only the Pangolin Node CA, so the TLS +// stack will verify any presented certificate against it. +// - VerifyPeerCertificate is wired to CRL.VerifyPeerCertificate: if the leaf cert +// CN is revoked the handshake fails immediately, before any RPC handler runs. +// - Minimum TLS version is 1.3. +// +// The caller must set cfg.Certificates with the server's own TLS certificate +// (typically from Let's Encrypt) before using this config. +func NewServerTLSConfig(ca *CA, crl *CRL) *tls.Config { + pool := x509.NewCertPool() + pool.AddCert(ca.CACert()) + + return &tls.Config{ + ClientAuth: tls.VerifyClientCertIfGiven, + ClientCAs: pool, + MinVersion: tls.VersionTLS13, + VerifyPeerCertificate: crl.VerifyPeerCertificate, + } +} diff --git a/server/setup.sh b/server/setup.sh new file mode 100644 index 0000000..b27b892 --- /dev/null +++ b/server/setup.sh @@ -0,0 +1,6 @@ +#!/usr/bin/env bash +# Run from the server/ directory to fetch deps and run tests. +# Usage: cd server && bash setup.sh +set -euo pipefail +go mod tidy +go test ./internal/mtls/...