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 <noreply@anthropic.com>
This commit is contained in:
wangjia
2026-06-13 01:44:21 +08:00
parent a642bf16a2
commit 73fd281bb5
11 changed files with 1012 additions and 0 deletions
+9
View File
@@ -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
)
+62
View File
@@ -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
}
+99
View File
@@ -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")
}
}
+192
View File
@@ -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)
}
+134
View File
@@ -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")
}
}
+83
View File
@@ -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
}
+107
View File
@@ -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)
}
}
+119
View File
@@ -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
}
+169
View File
@@ -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")
}
}
+32
View File
@@ -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,
}
}
+6
View File
@@ -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/...