Source file src/runtime/secret/secret_test.go

     1  // Copyright 2024 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  // these tests rely on inspecting freed memory, so they
     6  // can't be run under any of the memory validating modes.
     7  // TODO: figure out just which test violate which condition
     8  // and split this file out by individual test cases.
     9  // There could be some value to running some of these
    10  // under validation
    11  
    12  //go:build goexperiment.runtimesecret && (arm64 || amd64) && linux && !race && !asan && !msan
    13  
    14  package secret
    15  
    16  import (
    17  	"encoding/hex"
    18  	"runtime"
    19  	"strings"
    20  	"testing"
    21  	"time"
    22  	"unsafe"
    23  	"weak"
    24  )
    25  
    26  type secretType int64
    27  
    28  const secretValue = 0x53c237_53c237
    29  
    30  // S is a type that might have some secrets in it.
    31  type S [100]secretType
    32  
    33  // makeS makes an S with secrets in it.
    34  //
    35  //go:noinline
    36  func makeS() S {
    37  	// Note: noinline ensures this doesn't get inlined and
    38  	// completely optimized away.
    39  	var s S
    40  	for i := range s {
    41  		s[i] = secretValue
    42  	}
    43  	return s
    44  }
    45  
    46  // heapS allocates an S on the heap with secrets in it.
    47  //
    48  //go:noinline
    49  func heapS() *S {
    50  	// Note: noinline forces heap allocation
    51  	s := makeS()
    52  	return &s
    53  }
    54  
    55  // for the tiny allocator
    56  //
    57  //go:noinline
    58  func heapSTiny() *secretType {
    59  	s := new(secretType(secretValue))
    60  	return s
    61  }
    62  
    63  // Test that when we allocate inside secret.Do, the resulting
    64  // allocations are zeroed by the garbage collector when they
    65  // are freed.
    66  // See runtime/mheap.go:freeSpecial.
    67  func TestHeap(t *testing.T) {
    68  	var addr uintptr
    69  	var p weak.Pointer[S]
    70  	Do(func() {
    71  		sp := heapS()
    72  		addr = uintptr(unsafe.Pointer(sp))
    73  		p = weak.Make(sp)
    74  	})
    75  	waitCollected(t, p)
    76  
    77  	// Check that object got zeroed.
    78  	checkRangeForSecret(t, addr, addr+unsafe.Sizeof(S{}))
    79  	// Also check our stack, just because we can.
    80  	checkStackForSecret(t)
    81  }
    82  
    83  func TestHeapTiny(t *testing.T) {
    84  	var addr uintptr
    85  	var p weak.Pointer[secretType]
    86  	Do(func() {
    87  		sp := heapSTiny()
    88  		addr = uintptr(unsafe.Pointer(sp))
    89  		p = weak.Make(sp)
    90  	})
    91  	waitCollected(t, p)
    92  
    93  	// Check that object got zeroed.
    94  	checkRangeForSecret(t, addr, addr+unsafe.Sizeof(secretType(0)))
    95  	// Also check our stack, just because we can.
    96  	checkStackForSecret(t)
    97  }
    98  
    99  // Test that when we return from secret.Do, we zero the stack used
   100  // by the argument to secret.Do.
   101  // See runtime/secret.go:secret_dec.
   102  func TestStack(t *testing.T) {
   103  	checkStackForSecret(t) // if this fails, something is wrong with the test
   104  
   105  	Do(func() {
   106  		s := makeS()
   107  		use(&s)
   108  	})
   109  
   110  	checkStackForSecret(t)
   111  }
   112  
   113  //go:noinline
   114  func use(s *S) {
   115  	// Note: noinline prevents dead variable elimination.
   116  }
   117  
   118  // Test that when we copy a stack, we zero the old one.
   119  // See runtime/stack.go:copystack.
   120  func TestStackCopy(t *testing.T) {
   121  	checkStackForSecret(t) // if this fails, something is wrong with the test
   122  
   123  	var lo, hi uintptr
   124  	Do(func() {
   125  		// Put some secrets on the current stack frame.
   126  		s := makeS()
   127  		use(&s)
   128  		// Remember the current stack.
   129  		lo, hi = getStack()
   130  		// Use a lot more stack to force a stack copy.
   131  		growStack()
   132  	})
   133  	checkRangeForSecret(t, lo, hi) // pre-grow stack
   134  	checkStackForSecret(t)         // post-grow stack (just because we can)
   135  }
   136  
   137  func growStack() {
   138  	growStack1(1000)
   139  }
   140  func growStack1(n int) {
   141  	if n == 0 {
   142  		return
   143  	}
   144  	growStack1(n - 1)
   145  }
   146  
   147  func TestPanic(t *testing.T) {
   148  	checkStackForSecret(t) // if this fails, something is wrong with the test
   149  
   150  	defer func() {
   151  		checkStackForSecret(t)
   152  
   153  		p := recover()
   154  		if p == nil {
   155  			t.Errorf("panic squashed")
   156  			return
   157  		}
   158  		var e error
   159  		var ok bool
   160  		if e, ok = p.(error); !ok {
   161  			t.Errorf("panic not an error")
   162  		}
   163  		if !strings.Contains(e.Error(), "divide by zero") {
   164  			t.Errorf("panic not a divide by zero error: %s", e.Error())
   165  		}
   166  		var pcs [10]uintptr
   167  		n := runtime.Callers(0, pcs[:])
   168  		frames := runtime.CallersFrames(pcs[:n])
   169  		for {
   170  			frame, more := frames.Next()
   171  			if strings.Contains(frame.Function, "dividePanic") {
   172  				t.Errorf("secret function in traceback")
   173  			}
   174  			if !more {
   175  				break
   176  			}
   177  		}
   178  	}()
   179  	Do(dividePanic)
   180  }
   181  
   182  func dividePanic() {
   183  	s := makeS()
   184  	use(&s)
   185  	_ = 8 / zero
   186  }
   187  
   188  var zero int
   189  
   190  func TestGoExit(t *testing.T) {
   191  	checkStackForSecret(t) // if this fails, something is wrong with the test
   192  
   193  	c := make(chan uintptr, 2)
   194  
   195  	go func() {
   196  		// Run the test in a separate goroutine
   197  		defer func() {
   198  			// Tell original goroutine what our stack is
   199  			// so it can check it for secrets.
   200  			lo, hi := getStack()
   201  			c <- lo
   202  			c <- hi
   203  		}()
   204  		Do(func() {
   205  			s := makeS()
   206  			use(&s)
   207  			// there's an entire round-trip through the scheduler between here
   208  			// and when we are able to check if the registers are still dirtied, and we're
   209  			// not guaranteed to run on the same M. Make a best effort attempt anyway
   210  			loadRegisters(unsafe.Pointer(&s))
   211  			runtime.Goexit()
   212  		})
   213  		t.Errorf("goexit didn't happen")
   214  	}()
   215  	lo := <-c
   216  	hi := <-c
   217  	// We want to wait until the other goroutine has finished Goexiting and
   218  	// cleared its stack. There's no signal for that, so just wait a bit.
   219  	time.Sleep(1 * time.Millisecond)
   220  
   221  	checkRangeForSecret(t, lo, hi)
   222  
   223  	var spillArea [64]secretType
   224  	n := spillRegisters(unsafe.Pointer(&spillArea))
   225  	if n > unsafe.Sizeof(spillArea) {
   226  		t.Fatalf("spill area overrun %d\n", n)
   227  	}
   228  	for i, v := range spillArea {
   229  		if v == secretValue {
   230  			t.Errorf("secret found in spill slot %d", i)
   231  		}
   232  	}
   233  }
   234  
   235  func checkStackForSecret(t *testing.T) {
   236  	t.Helper()
   237  	lo, hi := getStack()
   238  	checkRangeForSecret(t, lo, hi)
   239  }
   240  func checkRangeForSecret(t *testing.T, lo, hi uintptr) {
   241  	t.Helper()
   242  	found := false
   243  	for p := lo; p < hi; p += unsafe.Sizeof(secretType(0)) {
   244  		v := *(*secretType)(unsafe.Pointer(p))
   245  		if v == secretValue {
   246  			found = true
   247  			t.Errorf("secret found in [%x,%x] at %x", lo, hi, p)
   248  		}
   249  	}
   250  	if found {
   251  		s := unsafe.Slice((*byte)(unsafe.Pointer(lo)), hi-lo)
   252  		t.Logf("%s", hex.Dump(s))
   253  	}
   254  }
   255  
   256  func waitCollected[P any](t *testing.T, ptr weak.Pointer[P]) {
   257  	t.Helper()
   258  	i := 0
   259  	for ptr.Value() != nil {
   260  		runtime.GC()
   261  		i++
   262  		// 20 seems like a decent number of times to try
   263  		if i > 20 {
   264  			t.Errorf("value was never collected")
   265  		}
   266  	}
   267  	t.Logf("number of cycles until collection: %d", i)
   268  }
   269  
   270  func TestRegisters(t *testing.T) {
   271  	Do(func() {
   272  		s := makeS()
   273  		loadRegisters(unsafe.Pointer(&s))
   274  	})
   275  	var spillArea [64]secretType
   276  	n := spillRegisters(unsafe.Pointer(&spillArea))
   277  	if n > unsafe.Sizeof(spillArea) {
   278  		t.Fatalf("spill area overrun %d\n", n)
   279  	}
   280  	for i, v := range spillArea {
   281  		if v == secretValue {
   282  			t.Errorf("secret found in spill slot %d", i)
   283  		}
   284  	}
   285  }
   286  
   287  func TestSecretInheritance(t *testing.T) {
   288  	ch := make(chan bool, 2)
   289  	Do(func() {
   290  		ch <- Enabled()
   291  		go func() {
   292  			ch <- Enabled()
   293  			close(ch)
   294  		}()
   295  	})
   296  	for enabled := range ch {
   297  		if !enabled {
   298  			t.Error("secret mode not enabled for child goroutine")
   299  		}
   300  	}
   301  }
   302  
   303  func TestSignalStacks(t *testing.T) {
   304  	Do(func() {
   305  		s := makeS()
   306  		loadRegisters(unsafe.Pointer(&s))
   307  		// cause a signal with our secret state to dirty
   308  		// at least one of the signal stacks
   309  		func() {
   310  			defer func() {
   311  				x := recover()
   312  				if x == nil {
   313  					panic("did not get panic")
   314  				}
   315  			}()
   316  			var p *int
   317  			*p = 20
   318  		}()
   319  	})
   320  	// signal stacks aren't cleared until after
   321  	// the next GC after secret.Do returns
   322  	runtime.GC()
   323  	stk := make([]stack, 0, 100)
   324  	stk = appendSignalStacks(stk)
   325  	for _, s := range stk {
   326  		checkRangeForSecret(t, s.lo, s.hi)
   327  	}
   328  }
   329  
   330  // hooks into the runtime
   331  func getStack() (uintptr, uintptr)
   332  
   333  // Stack is a copy of runtime.stack for testing export.
   334  // Fields must match.
   335  type stack struct {
   336  	lo uintptr
   337  	hi uintptr
   338  }
   339  
   340  func appendSignalStacks([]stack) []stack
   341  

View as plain text