73fd281bb5
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>
170 lines
5.3 KiB
Go
170 lines
5.3 KiB
Go
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")
|
|
}
|
|
}
|