1package scram
2
3import (
4 "crypto/ed25519"
5 cryptorand "crypto/rand"
6 "crypto/sha1"
7 "crypto/sha256"
8 "crypto/tls"
9 "crypto/x509"
10 "encoding/base64"
11 "errors"
12 "hash"
13 "math/big"
14 "net"
15 "testing"
16 "time"
17)
18
19func base64Decode(s string) []byte {
20 buf, err := base64.StdEncoding.DecodeString(s)
21 if err != nil {
22 panic("bad base64")
23 }
24 return buf
25}
26
27func tcheck(t *testing.T, err error, msg string) {
28 t.Helper()
29 if err != nil {
30 t.Fatalf("%s: %s", msg, err)
31 }
32}
33
34func TestSCRAMSHA1Server(t *testing.T) {
35 // Test vector from ../rfc/5802:496
36 salt := base64Decode("QSXCR+Q6sek8bf92")
37 saltedPassword, err := SaltPassword(sha1.New, "pencil", salt, 4096)
38 tcheck(t, err, "saltpassword")
39
40 server, err := NewServer(sha1.New, []byte("n,,n=user,r=fyko+d2lbbFgONRv9qkxdawL"), nil, false)
41 server.serverNonceOverride = "3rfcNHYJY1ZVvWVs7j"
42 tcheck(t, err, "newserver")
43 resp, err := server.ServerFirst(4096, salt)
44 tcheck(t, err, "server first")
45 if resp != "r=fyko+d2lbbFgONRv9qkxdawL3rfcNHYJY1ZVvWVs7j,s=QSXCR+Q6sek8bf92,i=4096" {
46 t.Fatalf("bad server first")
47 }
48 serverFinal, err := server.Finish([]byte("c=biws,r=fyko+d2lbbFgONRv9qkxdawL3rfcNHYJY1ZVvWVs7j,p=v0X8v3Bz2T0CJGbJQyF0X+HI4Ts="), saltedPassword)
49 tcheck(t, err, "finish")
50 if serverFinal != "v=rmF9pqV8S7suAoZWja4dJRkFsKQ=" {
51 t.Fatalf("bad server final")
52 }
53}
54
55func TestSCRAMSHA256Server(t *testing.T) {
56 // Test vector from ../rfc/7677:122
57 salt := base64Decode("W22ZaJ0SNY7soEsUEjb6gQ==")
58 saltedPassword, err := SaltPassword(sha256.New, "pencil", salt, 4096)
59 tcheck(t, err, "saltpassword")
60
61 server, err := NewServer(sha256.New, []byte("n,,n=user,r=rOprNGfwEbeRWgbNEkqO"), nil, false)
62 server.serverNonceOverride = "%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0"
63 tcheck(t, err, "newserver")
64 resp, err := server.ServerFirst(4096, salt)
65 tcheck(t, err, "server first")
66 if resp != "r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,s=W22ZaJ0SNY7soEsUEjb6gQ==,i=4096" {
67 t.Fatalf("bad server first")
68 }
69 serverFinal, err := server.Finish([]byte("c=biws,r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,p=dHzbZapWIk4jUhN+Ute9ytag9zjfMHgsqmmiz7AndVQ="), saltedPassword)
70 tcheck(t, err, "finish")
71 if serverFinal != "v=6rriTRBi23WpRR/wtup+mMhUZUn/dB5nLTJRsjl95G4=" {
72 t.Fatalf("bad server final")
73 }
74}
75
76// Bad attempt with wrong password.
77func TestScramServerBadPassword(t *testing.T) {
78 salt := base64Decode("W22ZaJ0SNY7soEsUEjb6gQ==")
79 saltedPassword, err := SaltPassword(sha256.New, "marker", salt, 4096)
80 tcheck(t, err, "saltpassword")
81
82 server, err := NewServer(sha256.New, []byte("n,,n=user,r=rOprNGfwEbeRWgbNEkqO"), nil, false)
83 server.serverNonceOverride = "%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0"
84 tcheck(t, err, "newserver")
85 _, err = server.ServerFirst(4096, salt)
86 tcheck(t, err, "server first")
87 _, err = server.Finish([]byte("c=biws,r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,p=dHzbZapWIk4jUhN+Ute9ytag9zjfMHgsqmmiz7AndVQ="), saltedPassword)
88 if !errors.Is(err, ErrInvalidProof) {
89 t.Fatalf("got %v, expected ErrInvalidProof", err)
90 }
91}
92
93// Bad attempt with different number of rounds.
94func TestScramServerBadIterations(t *testing.T) {
95 salt := base64Decode("W22ZaJ0SNY7soEsUEjb6gQ==")
96 saltedPassword, err := SaltPassword(sha256.New, "pencil", salt, 2048)
97 tcheck(t, err, "saltpassword")
98
99 server, err := NewServer(sha256.New, []byte("n,,n=user,r=rOprNGfwEbeRWgbNEkqO"), nil, false)
100 server.serverNonceOverride = "%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0"
101 tcheck(t, err, "newserver")
102 _, err = server.ServerFirst(4096, salt)
103 tcheck(t, err, "server first")
104 _, err = server.Finish([]byte("c=biws,r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,p=dHzbZapWIk4jUhN+Ute9ytag9zjfMHgsqmmiz7AndVQ="), saltedPassword)
105 if !errors.Is(err, ErrInvalidProof) {
106 t.Fatalf("got %v, expected ErrInvalidProof", err)
107 }
108}
109
110// Another attempt but with a randomly different nonce.
111func TestScramServerBad(t *testing.T) {
112 salt := base64Decode("W22ZaJ0SNY7soEsUEjb6gQ==")
113 saltedPassword, err := SaltPassword(sha256.New, "pencil", salt, 4096)
114 tcheck(t, err, "saltpassword")
115
116 server, err := NewServer(sha256.New, []byte("n,,n=user,r=rOprNGfwEbeRWgbNEkqO"), nil, false)
117 tcheck(t, err, "newserver")
118 _, err = server.ServerFirst(4096, salt)
119 tcheck(t, err, "server first")
120 _, err = server.Finish([]byte("c=biws,r="+server.nonce+",p=dHzbZapWIk4jUhN+Ute9ytag9zjfMHgsqmmiz7AndVQ="), saltedPassword)
121 if !errors.Is(err, ErrInvalidProof) {
122 t.Fatalf("got %v, expected ErrInvalidProof", err)
123 }
124}
125
126// Test with multiple extensions.
127func TestScramExtensions(t *testing.T) {
128 salt := base64Decode("W22ZaJ0SNY7soEsUEjb6gQ==")
129 saltedPassword, err := SaltPassword(sha256.New, "pencil", salt, 4096)
130 tcheck(t, err, "saltpassword")
131
132 // Client-first with two trailing extensions (a=1,b=2).
133 server, err := NewServer(sha256.New, []byte("n,,n=user,r=rOprNGfwEbeRWgbNEkqO,a=1,b=2"), nil, false)
134 tcheck(t, err, "newserver with client-first extensions")
135 if server.Authentication != "user" {
136 t.Fatalf("got username %q, expected %q", server.Authentication, "user")
137 }
138 if server.clientNonce != "rOprNGfwEbeRWgbNEkqO" {
139 t.Fatalf("got client nonce %q, expected %q", server.clientNonce, "rOprNGfwEbeRWgbNEkqO")
140 }
141
142 _, err = server.ServerFirst(4096, salt)
143 tcheck(t, err, "server first")
144
145 // Client-final with two extensions (a=1,b=2).
146 _, err = server.Finish([]byte("c=biws,r="+server.nonce+",a=1,b=2,p=dHzbZapWIk4jUhN+Ute9ytag9zjfMHgsqmmiz7AndVQ="), saltedPassword)
147 if !errors.Is(err, ErrInvalidProof) {
148 t.Fatalf("got %v, expected ErrInvalidProof", err)
149 }
150}
151
152func TestScramClient(t *testing.T) {
153 c := NewClient(sha256.New, "user", "", false, nil)
154 c.clientNonce = "rOprNGfwEbeRWgbNEkqO"
155 clientFirst, err := c.ClientFirst()
156 tcheck(t, err, "ClientFirst")
157 if clientFirst != "n,,n=user,r=rOprNGfwEbeRWgbNEkqO" {
158 t.Fatalf("bad clientFirst")
159 }
160 clientFinal, err := c.ServerFirst([]byte("r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,s=W22ZaJ0SNY7soEsUEjb6gQ==,i=4096"), "pencil")
161 tcheck(t, err, "ServerFirst")
162 if clientFinal != "c=biws,r=rOprNGfwEbeRWgbNEkqO%hvYDpWUa2RaTCAfuxFIlj)hNlF$k0,p=dHzbZapWIk4jUhN+Ute9ytag9zjfMHgsqmmiz7AndVQ=" {
163 t.Fatalf("bad clientFinal")
164 }
165 err = c.ServerFinal([]byte("v=6rriTRBi23WpRR/wtup+mMhUZUn/dB5nLTJRsjl95G4="))
166 tcheck(t, err, "ServerFinal")
167}
168
169func TestScram(t *testing.T) {
170 runHash := func(h func() hash.Hash, expErr error, username, authzid, password string, iterations int, clientNonce, serverNonce string, noServerPlus bool, clientcs, servercs *tls.ConnectionState) {
171 t.Helper()
172
173 defer func() {
174 x := recover()
175 if x == nil || x == "" {
176 return
177 }
178 panic(x)
179 }()
180
181 // check err is either nil or the expected error. if the expected error, panic to abort the authentication session.
182 xerr := func(err error, msg string) {
183 t.Helper()
184 if err != nil && !errors.Is(err, expErr) {
185 t.Fatalf("%s: got %v, expected %v", msg, err, expErr)
186 }
187 if err != nil {
188 panic("") // Abort test.
189 }
190 }
191
192 salt := MakeRandom()
193 saltedPassword, err := SaltPassword(h, password, salt, iterations)
194 tcheck(t, err, "saltpassword")
195
196 client := NewClient(h, username, "", noServerPlus, clientcs)
197 client.clientNonce = clientNonce
198 clientFirst, err := client.ClientFirst()
199 xerr(err, "client.ClientFirst")
200
201 server, err := NewServer(h, []byte(clientFirst), servercs, servercs != nil)
202 xerr(err, "NewServer")
203 server.serverNonceOverride = serverNonce
204
205 serverFirst, err := server.ServerFirst(iterations, salt)
206 xerr(err, "server.ServerFirst")
207
208 clientFinal, err := client.ServerFirst([]byte(serverFirst), password)
209 xerr(err, "client.ServerFirst")
210
211 serverFinal, err := server.Finish([]byte(clientFinal), saltedPassword)
212 xerr(err, "server.Finish")
213
214 err = client.ServerFinal([]byte(serverFinal))
215 xerr(err, "client.ServerFinal")
216
217 if expErr != nil {
218 t.Fatalf("got no error, expected %v", expErr)
219 }
220 }
221
222 makeState := func(maxTLSVersion uint16) (tls.ConnectionState, tls.ConnectionState) {
223 client, server := net.Pipe()
224 defer client.Close()
225 defer server.Close()
226 tlsClient := tls.Client(client, &tls.Config{
227 InsecureSkipVerify: true,
228 MaxVersion: maxTLSVersion,
229 })
230 tlsServer := tls.Server(server, &tls.Config{
231 Certificates: []tls.Certificate{fakeCert(t, "mox.example", false)},
232 MaxVersion: maxTLSVersion,
233 })
234 errc := make(chan error, 1)
235 go func() {
236 errc <- tlsServer.Handshake()
237 }()
238 err := tlsClient.Handshake()
239 tcheck(t, err, "tls handshake")
240 err = <-errc
241 tcheck(t, err, "server tls handshake")
242 clientcs := tlsClient.ConnectionState()
243 servercs := tlsServer.ConnectionState()
244
245 return clientcs, servercs
246 }
247
248 runPlus := func(maxTLSVersion uint16, expErr error, username, authzid, password string, iterations int, clientNonce, serverNonce string) {
249 t.Helper()
250
251 // PLUS variants.
252 clientcs, servercs := makeState(maxTLSVersion)
253 runHash(sha1.New, expErr, username, authzid, password, iterations, clientNonce, serverNonce, false, &clientcs, &servercs)
254 runHash(sha256.New, expErr, username, authzid, password, iterations, clientNonce, serverNonce, false, &clientcs, &servercs)
255 }
256
257 run := func(expErr error, username, authzid, password string, iterations int, clientNonce, serverNonce string) {
258 t.Helper()
259
260 // Bare variants
261 runHash(sha1.New, expErr, username, authzid, password, iterations, clientNonce, serverNonce, false, nil, nil)
262 runHash(sha256.New, expErr, username, authzid, password, iterations, clientNonce, serverNonce, false, nil, nil)
263
264 // Check with both TLS 1.2 for "tls-unique", and latest TLS for "tls-exporter".
265 runPlus(tls.VersionTLS12, expErr, username, authzid, password, iterations, clientNonce, serverNonce)
266 runPlus(0, expErr, username, authzid, password, iterations, clientNonce, serverNonce)
267 }
268
269 run(nil, "user", "", "pencil", 4096, "", "")
270 run(nil, "mjl@mox.example", "", "testtest", 4096, "", "")
271 run(nil, "mjl@mox.example", "", "short", 4096, "", "")
272 run(nil, "mjl@mox.example", "", "short", 2048, "", "")
273 run(nil, "mjl@mox.example", "mjl@mox.example", "testtest", 4096, "", "")
274 run(nil, "mjl@mox.example", "other@mox.example", "testtest", 4096, "", "")
275 run(nil, "mjl@mox.example", "other@mox.example", "testtest", 4096, "", "")
276 run(ErrUnsafe, "user", "", "pencil", 1, "", "") // Few iterations.
277 run(ErrUnsafe, "user", "", "pencil", 2048, "short", "") // Short client nonce.
278 run(ErrUnsafe, "user", "", "pencil", 2048, "test1234", "test") // Server added too few random data.
279
280 // Test mechanism downgrade attacks are detected.
281 runHash(sha1.New, ErrServerDoesSupportChannelBinding, "user", "", "pencil", 4096, "", "", true, nil, nil)
282 runHash(sha256.New, ErrServerDoesSupportChannelBinding, "user", "", "pencil", 4096, "", "", true, nil, nil)
283
284 // Test channel binding, detecting MitM attacks.
285 runChannelBind := func(maxTLSVersion uint16) {
286 t.Helper()
287
288 clientcs0, _ := makeState(maxTLSVersion)
289 _, servercs1 := makeState(maxTLSVersion)
290 runHash(sha1.New, ErrChannelBindingsDontMatch, "user", "", "pencil", 4096, "", "", false, &clientcs0, &servercs1)
291 runHash(sha256.New, ErrChannelBindingsDontMatch, "user", "", "pencil", 4096, "", "", false, &clientcs0, &servercs1)
292
293 // Client thinks it is on a TLS connection and server is not.
294 runHash(sha1.New, ErrChannelBindingsDontMatch, "user", "", "pencil", 4096, "", "", false, &clientcs0, nil)
295 runHash(sha256.New, ErrChannelBindingsDontMatch, "user", "", "pencil", 4096, "", "", false, &clientcs0, nil)
296 }
297
298 runChannelBind(0)
299 runChannelBind(tls.VersionTLS12)
300}
301
302// Just a cert that appears valid.
303func fakeCert(t *testing.T, name string, expired bool) tls.Certificate {
304 notAfter := time.Now()
305 if expired {
306 notAfter = notAfter.Add(-time.Hour)
307 } else {
308 notAfter = notAfter.Add(time.Hour)
309 }
310
311 privKey := ed25519.NewKeyFromSeed(make([]byte, ed25519.SeedSize)) // Fake key, don't use this for real!
312 template := &x509.Certificate{
313 SerialNumber: big.NewInt(1), // Required field...
314 DNSNames: []string{name},
315 NotBefore: time.Now().Add(-time.Hour),
316 NotAfter: notAfter,
317 }
318 localCertBuf, err := x509.CreateCertificate(cryptorand.Reader, template, template, privKey.Public(), privKey)
319 if err != nil {
320 t.Fatalf("making certificate: %s", err)
321 }
322 cert, err := x509.ParseCertificate(localCertBuf)
323 if err != nil {
324 t.Fatalf("parsing generated certificate: %s", err)
325 }
326 c := tls.Certificate{
327 Certificate: [][]byte{localCertBuf},
328 PrivateKey: privKey,
329 Leaf: cert,
330 }
331 return c
332}
333