v0.3.1 minor refactor, tests changed to standard library
This commit is contained in:
+661
-295
@@ -1,334 +1,700 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"math"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/crypto/argon2"
|
||||
)
|
||||
|
||||
// setupScramTest is a helper to initialize a server and user credential for testing.
|
||||
// It performs the full Argon2 -> SCRAM migration workflow.
|
||||
func setupScramTest(t *testing.T) (server *ScramServer, username, password string, cred *Credential) {
|
||||
username = "testuser"
|
||||
password = "SecurePassword123"
|
||||
|
||||
// 1. Start with an Argon2 PHC hash, as a real application would.
|
||||
phcHash, err := HashPassword(password)
|
||||
require.NoError(t, err, "Setup failed: could not hash password")
|
||||
|
||||
// 2. Migrate the PHC hash to a SCRAM credential.
|
||||
cred, err = MigrateFromPHC(username, password, phcHash)
|
||||
require.NoError(t, err, "Setup failed: could not migrate from PHC hash")
|
||||
|
||||
// 3. Create a server and add the new credential.
|
||||
server = NewScramServer()
|
||||
t.Cleanup(server.Stop)
|
||||
server.AddCredential(cred)
|
||||
|
||||
return server, username, password, cred
|
||||
func newTestServer(t *testing.T) *ScramServer {
|
||||
t.Helper()
|
||||
s := NewScramServer()
|
||||
t.Cleanup(s.Stop)
|
||||
return s
|
||||
}
|
||||
|
||||
// TestScram_FullRoundtrip_Success simulates a complete, successful authentication handshake.
|
||||
func TestScram_FullRoundtrip_Success(t *testing.T) {
|
||||
server, username, password, _ := setupScramTest(t)
|
||||
client := NewScramClient(username, password)
|
||||
|
||||
// --- Step 1: Client sends its first message ---
|
||||
clientFirst, err := client.StartAuthentication()
|
||||
require.NoError(t, err)
|
||||
|
||||
// --- Step 2: Server receives client's message and responds ---
|
||||
serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce)
|
||||
require.NoError(t, err, "Server failed to process client's first message")
|
||||
|
||||
// --- Step 3: Client receives server's message, computes proof ---
|
||||
clientFinal, err := client.ProcessServerFirstMessage(serverFirst)
|
||||
require.NoError(t, err, "Client failed to process server's first message")
|
||||
|
||||
// --- Step 4: Server receives client's proof and verifies it ---
|
||||
serverFinal, err := server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof)
|
||||
require.NoError(t, err, "Server failed to verify client's final proof")
|
||||
assert.NotEmpty(t, serverFinal.ServerSignature, "Server signature should not be empty")
|
||||
|
||||
// --- Step 5: Client verifies server's signature (mutual authentication) ---
|
||||
err = client.VerifyServerFinalMessage(serverFinal)
|
||||
assert.NoError(t, err, "Client failed to verify server's final signature")
|
||||
|
||||
t.Log("SCRAM full roundtrip successful")
|
||||
func testCredential(t *testing.T, username, password string) *Credential {
|
||||
t.Helper()
|
||||
phcHash, err := HashPassword(password, cheapArgon...)
|
||||
noErr(t, err, "HashPassword")
|
||||
cred, err := MigrateFromPHC(username, password, phcHash)
|
||||
noErr(t, err, "MigrateFromPHC")
|
||||
return cred
|
||||
}
|
||||
|
||||
// TestScram_FullRoundtrip_WrongPassword ensures authentication fails with an incorrect password.
|
||||
func TestScram_FullRoundtrip_WrongPassword(t *testing.T) {
|
||||
server, username, _, _ := setupScramTest(t)
|
||||
defer server.Stop()
|
||||
// Create a client with the WRONG password
|
||||
client := NewScramClient(username, "WrongPassword!!!")
|
||||
|
||||
// Steps 1-3 will appear to succeed, as the client doesn't know the password is wrong yet.
|
||||
clientFirst, err := client.StartAuthentication()
|
||||
require.NoError(t, err)
|
||||
|
||||
serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce)
|
||||
require.NoError(t, err)
|
||||
|
||||
clientFinal, err := client.ProcessServerFirstMessage(serverFirst)
|
||||
require.NoError(t, err)
|
||||
|
||||
// --- Step 4: Server verification should fail here ---
|
||||
_, err = server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof)
|
||||
assert.ErrorIs(t, err, ErrInvalidCredentials, "Server should reject proof from wrong password")
|
||||
|
||||
t.Log("SCRAM correctly failed for wrong password")
|
||||
func setupScram(t *testing.T) (*ScramServer, string, string, *Credential) {
|
||||
t.Helper()
|
||||
const username, password = "testuser", "SecurePassword123"
|
||||
cred := testCredential(t, username, password)
|
||||
s := newTestServer(t)
|
||||
s.AddCredential(cred)
|
||||
return s, username, password, cred
|
||||
}
|
||||
|
||||
// TestScram_FullRoundtrip_UserNotFound tests for user enumeration protection.
|
||||
// The server must be indistinguishable from the wrong-password path: no error at
|
||||
// first message, stable decoy salt across probes, ErrInvalidCredentials at proof.
|
||||
func TestScram_FullRoundtrip_UserNotFound(t *testing.T) {
|
||||
server, _, _, _ := setupScramTest(t)
|
||||
defer server.Stop()
|
||||
client := NewScramClient("unknown_user", "any_password")
|
||||
|
||||
clientFirst, err := client.StartAuthentication()
|
||||
require.NoError(t, err)
|
||||
|
||||
// unknown user must not error here
|
||||
serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce)
|
||||
require.NoError(t, err, "unknown user must not be signalled at first message")
|
||||
assert.NotEmpty(t, serverFirst.FullNonce)
|
||||
assert.NotEmpty(t, serverFirst.Salt)
|
||||
|
||||
// decoy salt must be stable across repeated probes
|
||||
second, err := server.ProcessClientFirstMessage(clientFirst.Username, "probe-nonce-2")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, serverFirst.Salt, second.Salt, "decoy salt must be deterministic")
|
||||
assert.Equal(t, serverFirst.ArgonTime, second.ArgonTime)
|
||||
assert.Equal(t, serverFirst.ArgonMemory, second.ArgonMemory)
|
||||
assert.Equal(t, serverFirst.ArgonThreads, second.ArgonThreads)
|
||||
|
||||
clientFinal, err := client.ProcessServerFirstMessage(serverFirst)
|
||||
require.NoError(t, err)
|
||||
|
||||
// failure surfaces only here, identical to wrong-password path
|
||||
_, err = server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof)
|
||||
assert.ErrorIs(t, err, ErrInvalidCredentials, "unknown user must fail like wrong password")
|
||||
func handshakeCount(s *ScramServer) int {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return len(s.handshakes)
|
||||
}
|
||||
|
||||
// TestScram_InvalidNonce simulates a replay attack or message mismatch.
|
||||
func TestScram_InvalidNonce(t *testing.T) {
|
||||
server, username, password, _ := setupScramTest(t)
|
||||
client := NewScramClient(username, password)
|
||||
|
||||
// Perform the first part of the handshake
|
||||
clientFirst, _ := client.StartAuthentication()
|
||||
serverFirst, _ := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce)
|
||||
clientFinal, _ := client.ProcessServerFirstMessage(serverFirst)
|
||||
|
||||
// Attempt to finalize with a completely different nonce
|
||||
_, err := server.ProcessClientFinalMessage("this-is-a-bad-nonce", clientFinal.ClientProof)
|
||||
assert.ErrorIs(t, err, ErrSCRAMInvalidNonce, "Server should reject a final message with an unknown nonce")
|
||||
// startHandshake drives a fresh client to the point where a proof is pending.
|
||||
func startHandshake(t *testing.T, s *ScramServer, username, password string) ClientFinalRequest {
|
||||
t.Helper()
|
||||
c := NewScramClient(username, password)
|
||||
first, err := c.StartAuthentication()
|
||||
noErr(t, err, "StartAuthentication")
|
||||
serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce)
|
||||
noErr(t, err, "ProcessClientFirstMessage")
|
||||
final, err := c.ProcessServerFirstMessage(serverFirst)
|
||||
noErr(t, err, "ProcessServerFirstMessage")
|
||||
return final
|
||||
}
|
||||
|
||||
// TestScram_CredentialImportExport verifies that credentials can be serialized and deserialized correctly.
|
||||
func TestScram_CredentialImportExport(t *testing.T) {
|
||||
_, _, _, originalCred := setupScramTest(t)
|
||||
|
||||
// Export the credential to a map
|
||||
exportedData := originalCred.Export()
|
||||
require.NotNil(t, exportedData)
|
||||
|
||||
// Assert that required fields exist and are strings (as they are base64 encoded)
|
||||
assert.IsType(t, "", exportedData["salt"])
|
||||
assert.IsType(t, "", exportedData["stored_key"])
|
||||
assert.IsType(t, "", exportedData["server_key"])
|
||||
|
||||
// Import the credential back from the map
|
||||
importedCred, err := ImportCredential(exportedData)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, importedCred)
|
||||
|
||||
// Verify that the imported credential is identical to the original
|
||||
assert.Equal(t, originalCred.Username, importedCred.Username)
|
||||
assert.Equal(t, originalCred.Salt, importedCred.Salt)
|
||||
assert.Equal(t, originalCred.ArgonTime, importedCred.ArgonTime)
|
||||
assert.Equal(t, originalCred.ArgonMemory, importedCred.ArgonMemory)
|
||||
assert.Equal(t, originalCred.ArgonThreads, importedCred.ArgonThreads)
|
||||
assert.Equal(t, originalCred.StoredKey, importedCred.StoredKey)
|
||||
assert.Equal(t, originalCred.ServerKey, importedCred.ServerKey)
|
||||
|
||||
t.Log("SCRAM credential import/export successful")
|
||||
}
|
||||
|
||||
// TestScramServerCleanup verifies automatic cleanup of expired handshakes
|
||||
func TestScramServerCleanup(t *testing.T) {
|
||||
// Create server with short cleanup interval for testing
|
||||
server := NewScramServer()
|
||||
defer server.Stop()
|
||||
|
||||
// Add a test credential
|
||||
cred := &Credential{
|
||||
Username: "testuser",
|
||||
Salt: []byte("salt1234567890123456"),
|
||||
ArgonTime: 1,
|
||||
ArgonMemory: 64,
|
||||
ArgonThreads: 1,
|
||||
StoredKey: []byte("stored_key_placeholder"),
|
||||
ServerKey: []byte("server_key_placeholder"),
|
||||
func runHandshake(s *ScramServer, c *ScramClient) error {
|
||||
first, err := c.StartAuthentication()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
server.AddCredential(cred)
|
||||
serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
final, err := c.ProcessServerFirstMessage(serverFirst)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
serverFinal, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return c.VerifyServerFinalMessage(serverFinal)
|
||||
}
|
||||
|
||||
// Start multiple handshakes
|
||||
var nonces []string
|
||||
for i := 0; i < 5; i++ {
|
||||
clientNonce := fmt.Sprintf("client-nonce-%d", i)
|
||||
msg, err := server.ProcessClientFirstMessage("testuser", clientNonce)
|
||||
require.NoError(t, err)
|
||||
nonces = append(nonces, msg.FullNonce)
|
||||
func TestScramRoundtrip(t *testing.T) {
|
||||
s, user, pw, _ := setupScram(t)
|
||||
noErr(t, runHandshake(s, NewScramClient(user, pw)), "handshake")
|
||||
eq(t, handshakeCount(s), 0, "handshake retained after success")
|
||||
}
|
||||
|
||||
func TestScramWrongPassword(t *testing.T) {
|
||||
s, user, _, _ := setupScram(t)
|
||||
errIs(t, runHandshake(s, NewScramClient(user, "WrongPassword!!!")), ErrInvalidCredentials, "wrong password")
|
||||
eq(t, handshakeCount(s), 0, "handshake retained after failure")
|
||||
}
|
||||
|
||||
func TestScramUnknownUser(t *testing.T) {
|
||||
s, _, _, cred := setupScram(t)
|
||||
c := NewScramClient("unknown_user", "any_password")
|
||||
|
||||
first, err := c.StartAuthentication()
|
||||
noErr(t, err, "StartAuthentication")
|
||||
|
||||
serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce)
|
||||
noErr(t, err, "unknown user must not be signalled at the first message")
|
||||
|
||||
// the decoy must mirror the registered parameter shape
|
||||
eq(t, serverFirst.ArgonTime, cred.ArgonTime, "decoy time")
|
||||
eq(t, serverFirst.ArgonMemory, cred.ArgonMemory, "decoy memory")
|
||||
eq(t, serverFirst.ArgonThreads, cred.ArgonThreads, "decoy threads")
|
||||
decoySalt, err := base64.StdEncoding.DecodeString(serverFirst.Salt)
|
||||
noErr(t, err, "decoy salt decode")
|
||||
eq(t, len(decoySalt), len(cred.Salt), "decoy salt length")
|
||||
|
||||
// stable across probes, distinct per username
|
||||
second, err := s.ProcessClientFirstMessage("unknown_user", "probe-2")
|
||||
noErr(t, err, "second probe")
|
||||
eq(t, second.Salt, serverFirst.Salt, "decoy salt must be deterministic")
|
||||
third, err := s.ProcessClientFirstMessage("other_unknown", "probe-3")
|
||||
noErr(t, err, "third probe")
|
||||
if third.Salt == serverFirst.Salt {
|
||||
t.Fatal("decoy salt is not username-bound")
|
||||
}
|
||||
|
||||
// Verify all handshakes exist
|
||||
server.mu.RLock()
|
||||
assert.Len(t, server.handshakes, 5)
|
||||
server.mu.RUnlock()
|
||||
// failure surfaces only at the proof step, as for a wrong password
|
||||
final, err := c.ProcessServerFirstMessage(serverFirst)
|
||||
noErr(t, err, "ProcessServerFirstMessage")
|
||||
_, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
errIs(t, err, ErrInvalidCredentials, "unknown user")
|
||||
}
|
||||
|
||||
// Manually set old timestamp for first 3 handshakes
|
||||
server.mu.Lock()
|
||||
oldTime := time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
count := 0
|
||||
for nonce := range server.handshakes {
|
||||
if count < 3 {
|
||||
server.handshakes[nonce].CreatedAt = oldTime
|
||||
count++
|
||||
func TestScramDecoySaltIsolation(t *testing.T) {
|
||||
// decoyKey is per-instance: a shared decoy salt would be a global oracle
|
||||
// for account existence across a cluster.
|
||||
a, b := newTestServer(t), newTestServer(t)
|
||||
first, err := a.ProcessClientFirstMessage("ghost", "n1")
|
||||
noErr(t, err, "server a")
|
||||
second, err := b.ProcessClientFirstMessage("ghost", "n2")
|
||||
noErr(t, err, "server b")
|
||||
if first.Salt == second.Salt {
|
||||
t.Fatal("decoy salt is identical across server instances")
|
||||
}
|
||||
|
||||
// with no credential registered the template is empty and defaults apply
|
||||
raw, err := base64.StdEncoding.DecodeString(first.Salt)
|
||||
noErr(t, err, "decode")
|
||||
eq(t, len(raw), DefaultArgonSaltLen, "fallback salt length")
|
||||
eq(t, first.ArgonTime, uint32(DefaultArgonTime), "fallback time")
|
||||
eq(t, first.ArgonMemory, uint32(DefaultArgonMemory), "fallback memory")
|
||||
eq(t, first.ArgonThreads, uint8(DefaultArgonThreads), "fallback threads")
|
||||
}
|
||||
|
||||
func TestScramDecoySaltMultiBlock(t *testing.T) {
|
||||
s := newTestServer(t)
|
||||
cred := testCredential(t, "u", "SecurePassword123")
|
||||
cred.Salt = make([]byte, 48) // exceeds one HMAC-SHA256 block
|
||||
s.AddCredential(cred)
|
||||
|
||||
msg, err := s.ProcessClientFirstMessage("ghost", "n")
|
||||
noErr(t, err, "first message")
|
||||
raw, err := base64.StdEncoding.DecodeString(msg.Salt)
|
||||
noErr(t, err, "decode")
|
||||
eq(t, len(raw), 48, "decoy salt length")
|
||||
}
|
||||
|
||||
func TestScramReplayAndUnknownNonce(t *testing.T) {
|
||||
s, user, pw, _ := setupScram(t)
|
||||
final := startHandshake(t, s, user, pw)
|
||||
|
||||
_, err := s.ProcessClientFinalMessage("this-is-a-bad-nonce", final.ClientProof)
|
||||
errIs(t, err, ErrSCRAMInvalidNonce, "unknown nonce")
|
||||
|
||||
_, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
noErr(t, err, "first proof")
|
||||
|
||||
_, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
errIs(t, err, ErrSCRAMInvalidNonce, "replayed proof")
|
||||
}
|
||||
|
||||
func TestScramProofBinding(t *testing.T) {
|
||||
// A proof commits to its own auth message; moving it to another live
|
||||
// handshake for the same user must fail.
|
||||
s, user, pw, _ := setupScram(t)
|
||||
a := startHandshake(t, s, user, pw)
|
||||
b := startHandshake(t, s, user, pw)
|
||||
|
||||
_, err := s.ProcessClientFinalMessage(a.FullNonce, b.ClientProof)
|
||||
errIs(t, err, ErrInvalidCredentials, "cross-handshake proof")
|
||||
|
||||
// the rejected attempt consumed handshake a but left b intact
|
||||
_, err = s.ProcessClientFinalMessage(a.FullNonce, a.ClientProof)
|
||||
errIs(t, err, ErrSCRAMInvalidNonce, "handshake a consumed")
|
||||
_, err = s.ProcessClientFinalMessage(b.FullNonce, b.ClientProof)
|
||||
noErr(t, err, "handshake b unaffected")
|
||||
}
|
||||
|
||||
func TestScramProofEncoding(t *testing.T) {
|
||||
s, user, pw, _ := setupScram(t)
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
proof string
|
||||
want error
|
||||
}{
|
||||
{"not base64", "!!!not base64!!!", ErrSCRAMInvalidProof},
|
||||
{"empty", "", ErrSCRAMInvalidProofLen},
|
||||
{"short", base64.StdEncoding.EncodeToString(make([]byte, 16)), ErrSCRAMInvalidProofLen},
|
||||
{"long", base64.StdEncoding.EncodeToString(make([]byte, sha256.Size+1)), ErrSCRAMInvalidProofLen},
|
||||
{"zeroed", base64.StdEncoding.EncodeToString(make([]byte, sha256.Size)), ErrInvalidCredentials},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
final := startHandshake(t, s, user, pw)
|
||||
_, err := s.ProcessClientFinalMessage(final.FullNonce, tc.proof)
|
||||
errIs(t, err, tc.want, tc.name)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestScramVerifyInProgress(t *testing.T) {
|
||||
s, user, pw, _ := setupScram(t)
|
||||
final := startHandshake(t, s, user, pw)
|
||||
|
||||
s.mu.RLock()
|
||||
state := s.handshakes[final.FullNonce]
|
||||
s.mu.RUnlock()
|
||||
if state == nil {
|
||||
t.Fatal("handshake not registered")
|
||||
}
|
||||
|
||||
state.verifying.Store(1)
|
||||
_, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
errIs(t, err, ErrSCRAMVerifyInProgress, "concurrent verification")
|
||||
|
||||
// the rejected attempt must not consume the handshake
|
||||
state.verifying.Store(0)
|
||||
_, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
noErr(t, err, "retry after release")
|
||||
}
|
||||
|
||||
func TestScramTimeouts(t *testing.T) {
|
||||
s, user, pw, _ := setupScram(t)
|
||||
final := startHandshake(t, s, user, pw)
|
||||
|
||||
s.mu.Lock()
|
||||
s.handshakes[final.FullNonce].CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
s.mu.Unlock()
|
||||
|
||||
_, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
errIs(t, err, ErrSCRAMTimeout, "server-side timeout")
|
||||
_, err = s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
errIs(t, err, ErrSCRAMInvalidNonce, "expired handshake consumed")
|
||||
|
||||
// client-side clock
|
||||
c := NewScramClient(user, pw)
|
||||
_, err = c.StartAuthentication()
|
||||
noErr(t, err, "StartAuthentication")
|
||||
c.startTime = time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
|
||||
_, err = c.ProcessServerFirstMessage(ServerFirstMessage{
|
||||
FullNonce: "n",
|
||||
Salt: base64.StdEncoding.EncodeToString(make([]byte, 16)),
|
||||
ArgonTime: testArgonTime,
|
||||
ArgonMemory: testArgonMemory, ArgonThreads: testArgonThreads,
|
||||
})
|
||||
errIs(t, err, ErrSCRAMTimeout, "client timeout on server-first")
|
||||
|
||||
c.authMessage = "seeded"
|
||||
c.serverKey = make([]byte, sha256.Size)
|
||||
errIs(t, c.VerifyServerFinalMessage(ServerFinalMessage{}), ErrSCRAMTimeout, "client timeout on server-final")
|
||||
}
|
||||
|
||||
func TestScramCleanup(t *testing.T) {
|
||||
s, user, _, _ := setupScram(t)
|
||||
for i := range 5 {
|
||||
_, err := s.ProcessClientFirstMessage(user, fmt.Sprintf("client-nonce-%d", i))
|
||||
noErr(t, err, "first message")
|
||||
}
|
||||
eq(t, handshakeCount(s), 5, "registered handshakes")
|
||||
|
||||
s.mu.Lock()
|
||||
aged := 0
|
||||
for _, state := range s.handshakes {
|
||||
if aged == 3 {
|
||||
break
|
||||
}
|
||||
state.CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
aged++
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
s.cleanupExpiredHandshakes()
|
||||
eq(t, handshakeCount(s), 2, "after sweep")
|
||||
|
||||
// a handshake under verification survives the sweep regardless of age
|
||||
s.mu.Lock()
|
||||
for _, state := range s.handshakes {
|
||||
state.CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
state.verifying.Store(1)
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
s.cleanupExpiredHandshakes()
|
||||
eq(t, handshakeCount(s), 2, "verifying handshakes must not be evicted")
|
||||
}
|
||||
|
||||
func TestScramHandshakeCap(t *testing.T) {
|
||||
s, user, _, _ := setupScram(t)
|
||||
for i := range ScramMaxHandshakes {
|
||||
_, err := s.ProcessClientFirstMessage(user, fmt.Sprintf("n-%d", i))
|
||||
noErr(t, err, "first message")
|
||||
}
|
||||
eq(t, handshakeCount(s), ScramMaxHandshakes, "at capacity")
|
||||
|
||||
_, err := s.ProcessClientFirstMessage(user, "overflow")
|
||||
errIs(t, err, ErrSCRAMTooManyHandshakes, "known user at capacity")
|
||||
|
||||
// the cap precedes credential lookup, so it is not an enumeration oracle
|
||||
_, err = s.ProcessClientFirstMessage("ghost", "overflow")
|
||||
errIs(t, err, ErrSCRAMTooManyHandshakes, "unknown user at capacity")
|
||||
|
||||
s.mu.Lock()
|
||||
for _, state := range s.handshakes {
|
||||
state.CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
}
|
||||
s.mu.Unlock()
|
||||
|
||||
_, err = s.ProcessClientFirstMessage(user, "after-sweep")
|
||||
noErr(t, err, "capacity reclaimed by the opportunistic sweep")
|
||||
eq(t, handshakeCount(s), 1, "all expired slots reclaimed")
|
||||
}
|
||||
|
||||
func TestScramNonceUniqueness(t *testing.T) {
|
||||
s, user, _, _ := setupScram(t)
|
||||
seen := make(map[string]struct{}, 256)
|
||||
for range 256 {
|
||||
// a fixed client nonce must not produce a fixed full nonce
|
||||
msg, err := s.ProcessClientFirstMessage(user, "fixed-client-nonce")
|
||||
noErr(t, err, "first message")
|
||||
if _, dup := seen[msg.FullNonce]; dup {
|
||||
t.Fatalf("duplicate full nonce: %s", msg.FullNonce)
|
||||
}
|
||||
seen[msg.FullNonce] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
func TestScramStopIdempotent(t *testing.T) {
|
||||
s := NewScramServer()
|
||||
s.Stop()
|
||||
s.Stop() // stopOnce must absorb the second close
|
||||
}
|
||||
|
||||
func TestScramClientState(t *testing.T) {
|
||||
s, user, pw, _ := setupScram(t)
|
||||
c := NewScramClient(user, pw)
|
||||
|
||||
errIs(t, c.VerifyServerFinalMessage(ServerFinalMessage{}), ErrSCRAMInvalidState, "unstarted client")
|
||||
|
||||
first, err := c.StartAuthentication()
|
||||
noErr(t, err, "StartAuthentication")
|
||||
serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce)
|
||||
noErr(t, err, "ProcessClientFirstMessage")
|
||||
final, err := c.ProcessServerFirstMessage(serverFirst)
|
||||
noErr(t, err, "ProcessServerFirstMessage")
|
||||
serverFinal, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof)
|
||||
noErr(t, err, "ProcessClientFinalMessage")
|
||||
|
||||
tampered := serverFinal
|
||||
tampered.ServerSignature = base64.StdEncoding.EncodeToString(make([]byte, sha256.Size))
|
||||
errIs(t, c.VerifyServerFinalMessage(tampered), ErrSCRAMServerAuthFailed, "forged signature")
|
||||
tampered.ServerSignature = "!!!"
|
||||
errIs(t, c.VerifyServerFinalMessage(tampered), ErrSCRAMServerAuthFailed, "malformed signature")
|
||||
noErr(t, c.VerifyServerFinalMessage(serverFinal), "valid signature")
|
||||
|
||||
c.Reset()
|
||||
errIs(t, c.VerifyServerFinalMessage(serverFinal), ErrSCRAMInvalidState, "after reset")
|
||||
next, err := c.StartAuthentication()
|
||||
noErr(t, err, "restart")
|
||||
if next.ClientNonce == first.ClientNonce {
|
||||
t.Fatal("client nonce reused after Reset")
|
||||
}
|
||||
}
|
||||
|
||||
func TestScramClientRejectsBadServerFirst(t *testing.T) {
|
||||
c := NewScramClient("u", "SecurePassword123")
|
||||
_, err := c.StartAuthentication()
|
||||
noErr(t, err, "StartAuthentication")
|
||||
|
||||
_, err = c.ProcessServerFirstMessage(ServerFirstMessage{
|
||||
FullNonce: "n", Salt: "!!!",
|
||||
ArgonTime: testArgonTime, ArgonMemory: testArgonMemory, ArgonThreads: testArgonThreads,
|
||||
})
|
||||
errIs(t, err, ErrSCRAMInvalidSalt, "salt encoding")
|
||||
|
||||
// ☢ no upper bound is applied to server-supplied cost parameters
|
||||
good := base64.StdEncoding.EncodeToString(make([]byte, 16))
|
||||
for _, msg := range []ServerFirstMessage{
|
||||
{FullNonce: "n", Salt: good, ArgonTime: 0, ArgonMemory: testArgonMemory, ArgonThreads: 1},
|
||||
{FullNonce: "n", Salt: good, ArgonTime: 1, ArgonMemory: 0, ArgonThreads: 1},
|
||||
{FullNonce: "n", Salt: good, ArgonTime: 1, ArgonMemory: testArgonMemory, ArgonThreads: 0},
|
||||
} {
|
||||
_, err = c.ProcessServerFirstMessage(msg)
|
||||
errIs(t, err, ErrSCRAMZeroParams, "zero parameter")
|
||||
}
|
||||
|
||||
// A hostile server cannot dictate an unbounded KDF
|
||||
_, err = c.ProcessServerFirstMessage(ServerFirstMessage{
|
||||
FullNonce: "n", Salt: good,
|
||||
ArgonTime: 1, ArgonMemory: MaxVerifyArgonMemory + 1, ArgonThreads: 1,
|
||||
})
|
||||
errIs(t, err, ErrSCRAMParamsTooLarge, "memory over ceiling")
|
||||
}
|
||||
|
||||
func TestScramClientOversizedPassword(t *testing.T) {
|
||||
c := NewScramClient("u", strings.Repeat("a", MaxPasswordLen+1))
|
||||
_, err := c.StartAuthentication()
|
||||
errIs(t, err, ErrPasswordTooLong, "oversized password rejected before the KDF")
|
||||
}
|
||||
|
||||
func TestServerFirstMessageMarshal(t *testing.T) {
|
||||
// the auth message binds this exact encoding; changes break every client
|
||||
msg := ServerFirstMessage{
|
||||
FullNonce: "abc", Salt: "c2FsdA==",
|
||||
ArgonTime: 3, ArgonMemory: 65536, ArgonThreads: 4,
|
||||
}
|
||||
eq(t, msg.Marshal(), "r=abc,s=c2FsdA==,t=3,m=65536,p=4", "marshal")
|
||||
}
|
||||
|
||||
func TestScramMigratedNonStandardDigest(t *testing.T) {
|
||||
// MigrateFromPHC falls back to DeriveCredential for digests other than 32
|
||||
// bytes; the resulting credential must still complete a handshake.
|
||||
const user, pw = "legacy", "SecurePassword123"
|
||||
cred, err := MigrateFromPHC(user, pw, phcFor(pw, []byte("0123456789abcdef"), 20))
|
||||
noErr(t, err, "MigrateFromPHC")
|
||||
eq(t, len(cred.StoredKey), sha256.Size, "stored key length")
|
||||
|
||||
s := newTestServer(t)
|
||||
s.AddCredential(cred)
|
||||
noErr(t, runHandshake(s, NewScramClient(user, pw)), "handshake with migrated credential")
|
||||
}
|
||||
|
||||
func TestDeriveCredential(t *testing.T) {
|
||||
const pw = "SecurePassword123"
|
||||
salt := make([]byte, 16)
|
||||
for i := range salt {
|
||||
salt[i] = byte(i)
|
||||
}
|
||||
|
||||
first, err := DeriveCredential("u", pw, salt, testArgonTime, testArgonMemory, testArgonThreads)
|
||||
noErr(t, err, "DeriveCredential")
|
||||
second, err := DeriveCredential("u", pw, salt, testArgonTime, testArgonMemory, testArgonThreads)
|
||||
noErr(t, err, "DeriveCredential repeat")
|
||||
eqBytes(t, first.StoredKey, second.StoredKey, "deterministic stored key")
|
||||
eqBytes(t, first.ServerKey, second.ServerKey, "deterministic server key")
|
||||
|
||||
salted := argon2.IDKey([]byte(pw), salt, testArgonTime, testArgonMemory, testArgonThreads, DefaultArgonKeyLen)
|
||||
want := sha256.Sum256(computeHMAC(salted, []byte("Client Key")))
|
||||
eqBytes(t, first.StoredKey, want[:], "stored key derivation")
|
||||
eqBytes(t, first.ServerKey, computeHMAC(salted, []byte("Server Key")), "server key derivation")
|
||||
if bytes.Equal(first.StoredKey, salted) || bytes.Equal(first.ServerKey, salted) {
|
||||
t.Fatal("credential exposes the salted password")
|
||||
}
|
||||
|
||||
// a different password must not collide
|
||||
other, err := DeriveCredential("u", pw+"x", salt, testArgonTime, testArgonMemory, testArgonThreads)
|
||||
noErr(t, err, "DeriveCredential other password")
|
||||
if bytes.Equal(first.StoredKey, other.StoredKey) {
|
||||
t.Fatal("stored key is independent of the password")
|
||||
}
|
||||
|
||||
_, err = DeriveCredential("u", pw, make([]byte, 15), testArgonTime, testArgonMemory, testArgonThreads)
|
||||
errIs(t, err, ErrSCRAMSaltTooShort, "short salt")
|
||||
|
||||
for _, p := range []struct {
|
||||
time, memory uint32
|
||||
threads uint8
|
||||
}{{0, testArgonMemory, 1}, {1, 0, 1}, {1, testArgonMemory, 0}} {
|
||||
_, err = DeriveCredential("u", pw, salt, p.time, p.memory, p.threads)
|
||||
errIs(t, err, ErrSCRAMZeroParams, "zero parameter")
|
||||
}
|
||||
|
||||
_, err = DeriveCredential("u", strings.Repeat("a", MaxPasswordLen+1), salt,
|
||||
testArgonTime, testArgonMemory, testArgonThreads)
|
||||
errIs(t, err, ErrPasswordTooLong, "oversized password")
|
||||
}
|
||||
|
||||
func TestCredentialExportImportRoundTrip(t *testing.T) {
|
||||
cred := testCredential(t, "roundtrip", "SecurePassword123")
|
||||
|
||||
imported, err := ImportCredential(cred.Export())
|
||||
noErr(t, err, "ImportCredential")
|
||||
eq(t, imported.Username, cred.Username, "username")
|
||||
eqBytes(t, imported.Salt, cred.Salt, "salt")
|
||||
eq(t, imported.ArgonTime, cred.ArgonTime, "time")
|
||||
eq(t, imported.ArgonMemory, cred.ArgonMemory, "memory")
|
||||
eq(t, imported.ArgonThreads, cred.ArgonThreads, "threads")
|
||||
eqBytes(t, imported.StoredKey, cred.StoredKey, "stored key")
|
||||
eqBytes(t, imported.ServerKey, cred.ServerKey, "server key")
|
||||
|
||||
// JSON transport converts every number to float64
|
||||
raw, err := json.Marshal(cred.Export())
|
||||
noErr(t, err, "marshal")
|
||||
var decoded map[string]any
|
||||
noErr(t, json.Unmarshal(raw, &decoded), "unmarshal")
|
||||
viaJSON, err := ImportCredential(decoded)
|
||||
noErr(t, err, "import via JSON")
|
||||
eq(t, viaJSON.ArgonMemory, cred.ArgonMemory, "memory via JSON")
|
||||
eqBytes(t, viaJSON.StoredKey, cred.StoredKey, "stored key via JSON")
|
||||
|
||||
// int-typed input, as produced by YAML and TOML decoders
|
||||
m := cred.Export()
|
||||
m["argon_time"] = int(cred.ArgonTime)
|
||||
m["argon_memory"] = int(cred.ArgonMemory)
|
||||
m["argon_threads"] = int(cred.ArgonThreads)
|
||||
viaInt, err := ImportCredential(m)
|
||||
noErr(t, err, "import from int-typed map")
|
||||
eq(t, viaInt.ArgonThreads, cred.ArgonThreads, "threads via int")
|
||||
|
||||
// an imported credential must still authenticate
|
||||
s := newTestServer(t)
|
||||
s.AddCredential(imported)
|
||||
noErr(t, runHandshake(s, NewScramClient("roundtrip", "SecurePassword123")), "handshake after import")
|
||||
}
|
||||
|
||||
func TestImportCredentialErrors(t *testing.T) {
|
||||
cred := testCredential(t, "u", "SecurePassword123")
|
||||
with := func(mutate func(map[string]any)) map[string]any {
|
||||
m := cred.Export()
|
||||
mutate(m)
|
||||
return m
|
||||
}
|
||||
b64 := func(n int) string { return base64.StdEncoding.EncodeToString(make([]byte, n)) }
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
data map[string]any
|
||||
want error
|
||||
}{
|
||||
{"empty map", map[string]any{}, ErrCredMissingUsername},
|
||||
{"username wrong type", with(func(m map[string]any) { m["username"] = 42 }), ErrCredMissingUsername},
|
||||
|
||||
{"salt missing", with(func(m map[string]any) { delete(m, "salt") }), ErrCredMissingSalt},
|
||||
{"salt not base64", with(func(m map[string]any) { m["salt"] = "!!!" }), ErrCredInvalidSalt},
|
||||
{"salt too short", with(func(m map[string]any) { m["salt"] = b64(15) }), ErrSCRAMSaltTooShort},
|
||||
|
||||
{"time missing", with(func(m map[string]any) { delete(m, "argon_time") }), ErrCredMissingTime},
|
||||
{"time wrong type", with(func(m map[string]any) { m["argon_time"] = "3" }), ErrCredInvalidType},
|
||||
{"time fractional", with(func(m map[string]any) { m["argon_time"] = 3.5 }), ErrCredInvalidType},
|
||||
{"time negative float", with(func(m map[string]any) { m["argon_time"] = float64(-1) }), ErrCredInvalidType},
|
||||
{"time float overflow", with(func(m map[string]any) { m["argon_time"] = float64(math.MaxUint32 + 1) }), ErrCredInvalidType},
|
||||
{"time negative int", with(func(m map[string]any) { m["argon_time"] = -1 }), ErrCredInvalidType},
|
||||
{"time zero", with(func(m map[string]any) { m["argon_time"] = uint32(0) }), ErrSCRAMZeroParams},
|
||||
|
||||
{"memory missing", with(func(m map[string]any) { delete(m, "argon_memory") }), ErrCredMissingMemory},
|
||||
{"memory zero", with(func(m map[string]any) { m["argon_memory"] = uint32(0) }), ErrSCRAMZeroParams},
|
||||
|
||||
{"threads missing", with(func(m map[string]any) { delete(m, "argon_threads") }), ErrCredMissingThreads},
|
||||
{"threads wrong type", with(func(m map[string]any) { m["argon_threads"] = "4" }), ErrCredInvalidType},
|
||||
{"threads overflow", with(func(m map[string]any) { m["argon_threads"] = float64(256) }), ErrCredInvalidType},
|
||||
{"threads negative", with(func(m map[string]any) { m["argon_threads"] = -1 }), ErrCredInvalidType},
|
||||
{"threads zero", with(func(m map[string]any) { m["argon_threads"] = uint8(0) }), ErrSCRAMZeroParams},
|
||||
|
||||
{"stored key missing", with(func(m map[string]any) { delete(m, "stored_key") }), ErrCredMissingStoredKey},
|
||||
{"stored key not base64", with(func(m map[string]any) { m["stored_key"] = "!!!" }), ErrCredInvalidStoredKey},
|
||||
{"stored key short", with(func(m map[string]any) { m["stored_key"] = b64(sha256.Size - 1) }), ErrCredInvalidStoredKey},
|
||||
|
||||
{"server key missing", with(func(m map[string]any) { delete(m, "server_key") }), ErrCredMissingServerKey},
|
||||
{"server key not base64", with(func(m map[string]any) { m["server_key"] = "!!!" }), ErrCredInvalidServerKey},
|
||||
{"server key short", with(func(m map[string]any) { m["server_key"] = b64(sha256.Size - 1) }), ErrCredInvalidServerKey},
|
||||
|
||||
{"stored key long", with(func(m map[string]any) { m["stored_key"] = b64(sha256.Size + 1) }), ErrCredInvalidStoredKey},
|
||||
{"server key long", with(func(m map[string]any) { m["server_key"] = b64(sha256.Size + 1) }), ErrCredInvalidServerKey},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got, err := ImportCredential(tc.data)
|
||||
errIs(t, err, tc.want, tc.name)
|
||||
if got != nil {
|
||||
t.Fatalf("%s: credential returned alongside error", tc.name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestScramConcurrentSameUser(t *testing.T) {
|
||||
s, user, pw, _ := setupScram(t)
|
||||
|
||||
const n = 12
|
||||
errs := make(chan error, n)
|
||||
var wg sync.WaitGroup
|
||||
for range n {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
errs <- runHandshake(s, NewScramClient(user, pw))
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
|
||||
for err := range errs {
|
||||
if err != nil {
|
||||
t.Errorf("concurrent handshake: %v", err)
|
||||
}
|
||||
}
|
||||
server.mu.Unlock()
|
||||
|
||||
// Trigger cleanup manually
|
||||
server.cleanupExpiredHandshakes()
|
||||
|
||||
// Verify only 2 handshakes remain
|
||||
server.mu.RLock()
|
||||
assert.Len(t, server.handshakes, 2, "Expired handshakes should be cleaned up")
|
||||
server.mu.RUnlock()
|
||||
eq(t, handshakeCount(s), 0, "handshakes leaked")
|
||||
}
|
||||
|
||||
// TestScramConcurrentSameUser verifies multiple concurrent authentications for same user
|
||||
func TestScramConcurrentSameUser(t *testing.T) {
|
||||
server, username, password, _ := setupScramTest(t)
|
||||
defer server.Stop()
|
||||
|
||||
// Number of concurrent authentication attempts
|
||||
numAttempts := 10
|
||||
results := make(chan error, numAttempts)
|
||||
func TestScramConcurrentMixedTraffic(t *testing.T) {
|
||||
s := newTestServer(t)
|
||||
creds := make([]*Credential, 4)
|
||||
for i := range creds {
|
||||
creds[i] = testCredential(t, fmt.Sprintf("user-%d", i), "SecurePassword123")
|
||||
}
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < numAttempts; i++ {
|
||||
for _, cred := range creds {
|
||||
wg.Add(1)
|
||||
go func(attempt int) {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
|
||||
// Each goroutine performs full authentication
|
||||
client := NewScramClient(username, password)
|
||||
|
||||
// Step 1: Client first
|
||||
clientFirst, err := client.StartAuthentication()
|
||||
if err != nil {
|
||||
results <- err
|
||||
return
|
||||
}
|
||||
|
||||
// Step 2: Server first
|
||||
serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce)
|
||||
if err != nil {
|
||||
results <- err
|
||||
return
|
||||
}
|
||||
|
||||
// Step 3: Client final
|
||||
clientFinal, err := client.ProcessServerFirstMessage(serverFirst)
|
||||
if err != nil {
|
||||
results <- err
|
||||
return
|
||||
}
|
||||
|
||||
// Step 4: Server final
|
||||
serverFinal, err := server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof)
|
||||
if err != nil {
|
||||
results <- err
|
||||
return
|
||||
}
|
||||
|
||||
// Step 5: Client verify
|
||||
err = client.VerifyServerFinalMessage(serverFinal)
|
||||
results <- err
|
||||
}(i)
|
||||
s.AddCredential(cred)
|
||||
}()
|
||||
}
|
||||
for i := range 16 {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
// registration races against lookup; both paths take s.mu
|
||||
_, _ = s.ProcessClientFirstMessage(fmt.Sprintf("user-%d", i%4), fmt.Sprintf("n-%d", i))
|
||||
_, _ = s.ProcessClientFirstMessage(fmt.Sprintf("ghost-%d", i), fmt.Sprintf("g-%d", i))
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
close(results)
|
||||
}
|
||||
|
||||
// Verify all attempts succeeded
|
||||
successCount := 0
|
||||
for err := range results {
|
||||
if err == nil {
|
||||
successCount++
|
||||
} else {
|
||||
t.Logf("Auth attempt failed: %v", err)
|
||||
func FuzzImportCredential(f *testing.F) {
|
||||
cred, err := DeriveCredential("u", "SecurePassword123", make([]byte, 16),
|
||||
testArgonTime, testArgonMemory, testArgonThreads)
|
||||
if err != nil {
|
||||
f.Fatal(err)
|
||||
}
|
||||
seed, err := json.Marshal(cred.Export())
|
||||
if err != nil {
|
||||
f.Fatal(err)
|
||||
}
|
||||
f.Add(seed)
|
||||
f.Add([]byte(`{}`))
|
||||
f.Add([]byte(`{"username":"u","salt":"","argon_time":1e309}`))
|
||||
|
||||
f.Fuzz(func(t *testing.T, data []byte) {
|
||||
var m map[string]any
|
||||
if err := json.Unmarshal(data, &m); err != nil || m == nil {
|
||||
return
|
||||
}
|
||||
got, err := ImportCredential(m)
|
||||
if err != nil {
|
||||
if got != nil {
|
||||
t.Fatal("credential returned alongside error")
|
||||
}
|
||||
return
|
||||
}
|
||||
if len(got.Salt) < 16 {
|
||||
t.Fatalf("accepted salt of %d bytes", len(got.Salt))
|
||||
}
|
||||
if len(got.StoredKey) != sha256.Size || len(got.ServerKey) != sha256.Size {
|
||||
t.Fatalf("accepted keys of %d/%d bytes", len(got.StoredKey), len(got.ServerKey))
|
||||
}
|
||||
if got.ArgonTime == 0 || got.ArgonMemory == 0 || got.ArgonThreads == 0 {
|
||||
t.Fatalf("accepted zero parameters: %+v", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func BenchmarkScramHandshake(b *testing.B) {
|
||||
const user, pw = "bench", "SecurePassword123"
|
||||
cred, err := DeriveCredential(user, pw, make([]byte, 16), testArgonTime, testArgonMemory, testArgonThreads)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
s := NewScramServer()
|
||||
defer s.Stop()
|
||||
s.AddCredential(cred)
|
||||
|
||||
for b.Loop() {
|
||||
c := NewScramClient(user, pw)
|
||||
first, err := c.StartAuthentication()
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
serverFirst, err := s.ProcessClientFirstMessage(first.Username, first.ClientNonce)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
final, err := c.ProcessServerFirstMessage(serverFirst)
|
||||
if err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
if _, err := s.ProcessClientFinalMessage(final.FullNonce, final.ClientProof); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
assert.Equal(t, numAttempts, successCount,
|
||||
"All concurrent authentication attempts should succeed")
|
||||
|
||||
// Verify no handshakes remain after completion
|
||||
server.mu.RLock()
|
||||
assert.Empty(t, server.handshakes, "All handshakes should be cleaned up after completion")
|
||||
server.mu.RUnlock()
|
||||
}
|
||||
|
||||
// TestScramExplicitTimeout verifies timeout enforcement
|
||||
func TestScramExplicitTimeout(t *testing.T) {
|
||||
// Save original timeout and set shorter one for testing
|
||||
originalTimeout := ScramHandshakeTimeout
|
||||
// Note: Can't modify const at runtime, so we test with delay instead
|
||||
|
||||
server, username, password, _ := setupScramTest(t)
|
||||
defer server.Stop()
|
||||
|
||||
client := NewScramClient(username, password)
|
||||
|
||||
// Start authentication
|
||||
clientFirst, err := client.StartAuthentication()
|
||||
require.NoError(t, err)
|
||||
|
||||
serverFirst, err := server.ProcessClientFirstMessage(clientFirst.Username, clientFirst.ClientNonce)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Manually expire the handshake
|
||||
server.mu.Lock()
|
||||
for nonce := range server.handshakes {
|
||||
server.handshakes[nonce].CreatedAt = time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
}
|
||||
server.mu.Unlock()
|
||||
|
||||
// Client processes server message (should work, client tracks own timeout)
|
||||
clientFinal, err := client.ProcessServerFirstMessage(serverFirst)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Server should reject due to timeout
|
||||
_, err = server.ProcessClientFinalMessage(clientFinal.FullNonce, clientFinal.ClientProof)
|
||||
assert.ErrorIs(t, err, ErrSCRAMTimeout, "Server should reject expired handshake")
|
||||
|
||||
// Test client-side timeout
|
||||
client2 := NewScramClient(username, password)
|
||||
client2.startTime = time.Now().Add(-2 * ScramHandshakeTimeout)
|
||||
|
||||
_, err = client2.ProcessServerFirstMessage(serverFirst)
|
||||
assert.ErrorIs(t, err, ErrSCRAMTimeout, "Client should reject after timeout")
|
||||
|
||||
_ = originalTimeout // Suppress unused variable warning
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user