1
2
3
4
5
6
7
8
9
10
11
12
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
31 type S [100]secretType
32
33
34
35
36 func makeS() S {
37
38
39 var s S
40 for i := range s {
41 s[i] = secretValue
42 }
43 return s
44 }
45
46
47
48
49 func heapS() *S {
50
51 s := makeS()
52 return &s
53 }
54
55
56
57
58 func heapSTiny() *secretType {
59 s := new(secretType(secretValue))
60 return s
61 }
62
63
64
65
66
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
78 checkRangeForSecret(t, addr, addr+unsafe.Sizeof(S{}))
79
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
94 checkRangeForSecret(t, addr, addr+unsafe.Sizeof(secretType(0)))
95
96 checkStackForSecret(t)
97 }
98
99
100
101
102 func TestStack(t *testing.T) {
103 checkStackForSecret(t)
104
105 Do(func() {
106 s := makeS()
107 use(&s)
108 })
109
110 checkStackForSecret(t)
111 }
112
113
114 func use(s *S) {
115
116 }
117
118
119
120 func TestStackCopy(t *testing.T) {
121 checkStackForSecret(t)
122
123 var lo, hi uintptr
124 Do(func() {
125
126 s := makeS()
127 use(&s)
128
129 lo, hi = getStack()
130
131 growStack()
132 })
133 checkRangeForSecret(t, lo, hi)
134 checkStackForSecret(t)
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)
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)
192
193 c := make(chan uintptr, 2)
194
195 go func() {
196
197 defer func() {
198
199
200 lo, hi := getStack()
201 c <- lo
202 c <- hi
203 }()
204 Do(func() {
205 s := makeS()
206 use(&s)
207
208
209
210 loadRegisters(unsafe.Pointer(&s))
211 runtime.Goexit()
212 })
213 t.Errorf("goexit didn't happen")
214 }()
215 lo := <-c
216 hi := <-c
217
218
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
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
308
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
321
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
331 func getStack() (uintptr, uintptr)
332
333
334
335 type stack struct {
336 lo uintptr
337 hi uintptr
338 }
339
340 func appendSignalStacks([]stack) []stack
341
View as plain text