mirror of
https://github.com/sonr-io/crypto.git
synced 2026-08-02 15:31:38 +00:00
No commit suggestions generated
This commit is contained in:
+213
@@ -0,0 +1,213 @@
|
||||
// Package argon2 provides secure key derivation using Argon2id
|
||||
package argon2
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/argon2"
|
||||
)
|
||||
|
||||
// Config defines Argon2id parameters
|
||||
type Config struct {
|
||||
Time uint32 // Number of iterations
|
||||
Memory uint32 // Memory in KB
|
||||
Parallelism uint8 // Number of threads
|
||||
SaltLength uint32 // Salt length in bytes
|
||||
KeyLength uint32 // Output key length in bytes
|
||||
}
|
||||
|
||||
// DefaultConfig returns secure default parameters
|
||||
func DefaultConfig() *Config {
|
||||
return &Config{
|
||||
Time: 1,
|
||||
Memory: 64 * 1024, // 64MB
|
||||
Parallelism: 4,
|
||||
SaltLength: 32,
|
||||
KeyLength: 32,
|
||||
}
|
||||
}
|
||||
|
||||
// LightConfig returns lighter parameters for testing
|
||||
func LightConfig() *Config {
|
||||
return &Config{
|
||||
Time: 1,
|
||||
Memory: 16 * 1024, // 16MB
|
||||
Parallelism: 2,
|
||||
SaltLength: 16,
|
||||
KeyLength: 32,
|
||||
}
|
||||
}
|
||||
|
||||
// HighSecurityConfig returns high-security parameters
|
||||
func HighSecurityConfig() *Config {
|
||||
return &Config{
|
||||
Time: 3,
|
||||
Memory: 128 * 1024, // 128MB
|
||||
Parallelism: 4,
|
||||
SaltLength: 32,
|
||||
KeyLength: 32,
|
||||
}
|
||||
}
|
||||
|
||||
// KDF implements Argon2id key derivation
|
||||
type KDF struct {
|
||||
config *Config
|
||||
}
|
||||
|
||||
// New creates a new Argon2id KDF with the given configuration
|
||||
func New(config *Config) *KDF {
|
||||
if config == nil {
|
||||
config = DefaultConfig()
|
||||
}
|
||||
return &KDF{config: config}
|
||||
}
|
||||
|
||||
// DeriveKey derives a key from password and salt
|
||||
func (k *KDF) DeriveKey(password []byte, salt []byte) []byte {
|
||||
return argon2.IDKey(
|
||||
password,
|
||||
salt,
|
||||
k.config.Time,
|
||||
k.config.Memory,
|
||||
k.config.Parallelism,
|
||||
k.config.KeyLength,
|
||||
)
|
||||
}
|
||||
|
||||
// GenerateSalt generates a cryptographically secure salt
|
||||
func (k *KDF) GenerateSalt() ([]byte, error) {
|
||||
salt := make([]byte, k.config.SaltLength)
|
||||
if _, err := rand.Read(salt); err != nil {
|
||||
return nil, fmt.Errorf("failed to generate salt: %w", err)
|
||||
}
|
||||
return salt, nil
|
||||
}
|
||||
|
||||
// HashPassword generates a hash with embedded salt and parameters
|
||||
func (k *KDF) HashPassword(password []byte) (string, error) {
|
||||
salt, err := k.GenerateSalt()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
hash := k.DeriveKey(password, salt)
|
||||
|
||||
// Encode in PHC format: $argon2id$v=19$m=65536,t=1,p=4$salt$hash
|
||||
encodedSalt := base64.RawStdEncoding.EncodeToString(salt)
|
||||
encodedHash := base64.RawStdEncoding.EncodeToString(hash)
|
||||
|
||||
return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s",
|
||||
argon2.Version,
|
||||
k.config.Memory,
|
||||
k.config.Time,
|
||||
k.config.Parallelism,
|
||||
encodedSalt,
|
||||
encodedHash,
|
||||
), nil
|
||||
}
|
||||
|
||||
// VerifyPassword verifies a password against a PHC-formatted hash
|
||||
func VerifyPassword(password []byte, encodedHash string) (bool, error) {
|
||||
params, salt, hash, err := decodeHash(encodedHash)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
kdf := &KDF{config: params}
|
||||
derivedHash := kdf.DeriveKey(password, salt)
|
||||
|
||||
// Constant-time comparison
|
||||
return subtle.ConstantTimeCompare(hash, derivedHash) == 1, nil
|
||||
}
|
||||
|
||||
// decodeHash parses PHC-formatted Argon2id hash
|
||||
func decodeHash(encodedHash string) (*Config, []byte, []byte, error) {
|
||||
parts := strings.Split(encodedHash, "$")
|
||||
if len(parts) != 6 {
|
||||
return nil, nil, nil, fmt.Errorf("invalid hash format")
|
||||
}
|
||||
|
||||
if parts[1] != "argon2id" {
|
||||
return nil, nil, nil, fmt.Errorf("unsupported algorithm: %s", parts[1])
|
||||
}
|
||||
|
||||
var version int
|
||||
_, err := fmt.Sscanf(parts[2], "v=%d", &version)
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("failed to parse version: %w", err)
|
||||
}
|
||||
|
||||
if version != argon2.Version {
|
||||
return nil, nil, nil, fmt.Errorf("unsupported Argon2 version: %d", version)
|
||||
}
|
||||
|
||||
var memory, time uint32
|
||||
var parallelism uint8
|
||||
_, err = fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &memory, &time, ¶llelism)
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("failed to parse parameters: %w", err)
|
||||
}
|
||||
|
||||
salt, err := base64.RawStdEncoding.DecodeString(parts[4])
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("failed to decode salt: %w", err)
|
||||
}
|
||||
|
||||
hash, err := base64.RawStdEncoding.DecodeString(parts[5])
|
||||
if err != nil {
|
||||
return nil, nil, nil, fmt.Errorf("failed to decode hash: %w", err)
|
||||
}
|
||||
|
||||
config := &Config{
|
||||
Time: time,
|
||||
Memory: memory,
|
||||
Parallelism: parallelism,
|
||||
SaltLength: uint32(len(salt)),
|
||||
KeyLength: uint32(len(hash)),
|
||||
}
|
||||
|
||||
return config, salt, hash, nil
|
||||
}
|
||||
|
||||
// CompareHashes performs constant-time comparison of two hashes
|
||||
func CompareHashes(hash1, hash2 []byte) bool {
|
||||
return subtle.ConstantTimeCompare(hash1, hash2) == 1
|
||||
}
|
||||
|
||||
// ValidateConfig validates Argon2id parameters
|
||||
func ValidateConfig(config *Config) error {
|
||||
if config.Time < 1 {
|
||||
return fmt.Errorf("time must be at least 1")
|
||||
}
|
||||
if config.Memory < 8*1024 {
|
||||
return fmt.Errorf("memory must be at least 8MB")
|
||||
}
|
||||
if config.Parallelism < 1 {
|
||||
return fmt.Errorf("parallelism must be at least 1")
|
||||
}
|
||||
if config.SaltLength < 8 {
|
||||
return fmt.Errorf("salt length must be at least 8 bytes")
|
||||
}
|
||||
if config.KeyLength < 16 {
|
||||
return fmt.Errorf("key length must be at least 16 bytes")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// EstimateTime estimates the time required for key derivation
|
||||
func EstimateTime(config *Config, iterations int) string {
|
||||
// This is a rough estimate - actual time depends on hardware
|
||||
baseTime := float64(config.Time) * float64(config.Memory) / (64 * 1024)
|
||||
totalTime := baseTime * float64(iterations)
|
||||
|
||||
if totalTime < 1 {
|
||||
return fmt.Sprintf("%.2f ms", totalTime*1000)
|
||||
} else if totalTime < 60 {
|
||||
return fmt.Sprintf("%.2f s", totalTime)
|
||||
}
|
||||
return fmt.Sprintf("%.2f min", totalTime/60)
|
||||
}
|
||||
@@ -0,0 +1,396 @@
|
||||
package argon2
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestKDF_DeriveKey(t *testing.T) {
|
||||
kdf := New(DefaultConfig())
|
||||
|
||||
password := []byte("test-password")
|
||||
salt := []byte("salt-must-be-at-least-16-bytes!!")
|
||||
|
||||
// Derive key
|
||||
key := kdf.DeriveKey(password, salt)
|
||||
assert.Len(t, key, int(kdf.config.KeyLength))
|
||||
|
||||
// Same inputs should produce same key
|
||||
key2 := kdf.DeriveKey(password, salt)
|
||||
assert.Equal(t, key, key2)
|
||||
|
||||
// Different password should produce different key
|
||||
key3 := kdf.DeriveKey([]byte("different"), salt)
|
||||
assert.NotEqual(t, key, key3)
|
||||
|
||||
// Different salt should produce different key
|
||||
salt2 := []byte("different-salt-at-least-16-bytes")
|
||||
key4 := kdf.DeriveKey(password, salt2)
|
||||
assert.NotEqual(t, key, key4)
|
||||
}
|
||||
|
||||
func TestKDF_GenerateSalt(t *testing.T) {
|
||||
kdf := New(DefaultConfig())
|
||||
|
||||
salt1, err := kdf.GenerateSalt()
|
||||
require.NoError(t, err)
|
||||
assert.Len(t, salt1, int(kdf.config.SaltLength))
|
||||
|
||||
// Should generate different salt each time
|
||||
salt2, err := kdf.GenerateSalt()
|
||||
require.NoError(t, err)
|
||||
assert.NotEqual(t, salt1, salt2)
|
||||
}
|
||||
|
||||
func TestKDF_HashPassword(t *testing.T) {
|
||||
kdf := New(DefaultConfig())
|
||||
password := []byte("MySecureP@ssw0rd")
|
||||
|
||||
hash, err := kdf.HashPassword(password)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Check format
|
||||
assert.True(t, strings.HasPrefix(hash, "$argon2id$"))
|
||||
parts := strings.Split(hash, "$")
|
||||
assert.Len(t, parts, 6)
|
||||
|
||||
// Verify password
|
||||
valid, err := VerifyPassword(password, hash)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, valid)
|
||||
|
||||
// Wrong password should fail
|
||||
valid, err = VerifyPassword([]byte("wrong"), hash)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, valid)
|
||||
}
|
||||
|
||||
func TestVerifyPassword(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
password string
|
||||
hash string
|
||||
valid bool
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid password",
|
||||
password: "password123",
|
||||
hash: "$argon2id$v=19$m=65536,t=1,p=4$c2FsdC1tdXN0LWJlLWF0LWxlYXN0LTE2LWJ5dGVzISE$+4smaTt/N7ivKLrqsPIbTplUxDBRMxTKCYOcXWTJOEI",
|
||||
valid: true,
|
||||
},
|
||||
{
|
||||
name: "invalid password",
|
||||
password: "wrongpassword",
|
||||
hash: "$argon2id$v=19$m=65536,t=1,p=4$c2FsdC1tdXN0LWJlLWF0LWxlYXN0LTE2LWJ5dGVzISE$+4smaTt/N7ivKLrqsPIbTplUxDBRMxTKCYOcXWTJOEI",
|
||||
valid: false,
|
||||
},
|
||||
{
|
||||
name: "invalid format",
|
||||
password: "password",
|
||||
hash: "invalid-hash-format",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "wrong algorithm",
|
||||
password: "password",
|
||||
hash: "$bcrypt$v=19$m=65536,t=1,p=4$salt$hash",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
valid, err := VerifyPassword([]byte(tc.password), tc.hash)
|
||||
if tc.wantErr {
|
||||
assert.Error(t, err)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, tc.valid, valid)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigurations(t *testing.T) {
|
||||
configs := map[string]*Config{
|
||||
"default": DefaultConfig(),
|
||||
"light": LightConfig(),
|
||||
"high": HighSecurityConfig(),
|
||||
}
|
||||
|
||||
for name, config := range configs {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
err := ValidateConfig(config)
|
||||
assert.NoError(t, err)
|
||||
|
||||
kdf := New(config)
|
||||
password := []byte("test-password")
|
||||
|
||||
hash, err := kdf.HashPassword(password)
|
||||
require.NoError(t, err)
|
||||
|
||||
valid, err := VerifyPassword(password, hash)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, valid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateConfig(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
config *Config
|
||||
wantErr bool
|
||||
errMsg string
|
||||
}{
|
||||
{
|
||||
name: "valid config",
|
||||
config: DefaultConfig(),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "time too low",
|
||||
config: &Config{
|
||||
Time: 0,
|
||||
Memory: 64 * 1024,
|
||||
Parallelism: 4,
|
||||
SaltLength: 32,
|
||||
KeyLength: 32,
|
||||
},
|
||||
wantErr: true,
|
||||
errMsg: "time must be at least 1",
|
||||
},
|
||||
{
|
||||
name: "memory too low",
|
||||
config: &Config{
|
||||
Time: 1,
|
||||
Memory: 4 * 1024,
|
||||
Parallelism: 4,
|
||||
SaltLength: 32,
|
||||
KeyLength: 32,
|
||||
},
|
||||
wantErr: true,
|
||||
errMsg: "memory must be at least 8MB",
|
||||
},
|
||||
{
|
||||
name: "parallelism too low",
|
||||
config: &Config{
|
||||
Time: 1,
|
||||
Memory: 64 * 1024,
|
||||
Parallelism: 0,
|
||||
SaltLength: 32,
|
||||
KeyLength: 32,
|
||||
},
|
||||
wantErr: true,
|
||||
errMsg: "parallelism must be at least 1",
|
||||
},
|
||||
{
|
||||
name: "salt too short",
|
||||
config: &Config{
|
||||
Time: 1,
|
||||
Memory: 64 * 1024,
|
||||
Parallelism: 4,
|
||||
SaltLength: 4,
|
||||
KeyLength: 32,
|
||||
},
|
||||
wantErr: true,
|
||||
errMsg: "salt length must be at least 8 bytes",
|
||||
},
|
||||
{
|
||||
name: "key too short",
|
||||
config: &Config{
|
||||
Time: 1,
|
||||
Memory: 64 * 1024,
|
||||
Parallelism: 4,
|
||||
SaltLength: 32,
|
||||
KeyLength: 8,
|
||||
},
|
||||
wantErr: true,
|
||||
errMsg: "key length must be at least 16 bytes",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
err := ValidateConfig(tc.config)
|
||||
if tc.wantErr {
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), tc.errMsg)
|
||||
} else {
|
||||
assert.NoError(t, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompareHashes(t *testing.T) {
|
||||
hash1 := []byte("hash1")
|
||||
hash2 := []byte("hash1")
|
||||
hash3 := []byte("hash2")
|
||||
|
||||
assert.True(t, CompareHashes(hash1, hash2))
|
||||
assert.False(t, CompareHashes(hash1, hash3))
|
||||
assert.False(t, CompareHashes([]byte("short"), []byte("longer")))
|
||||
}
|
||||
|
||||
func TestEstimateTime(t *testing.T) {
|
||||
config := DefaultConfig()
|
||||
|
||||
estimate := EstimateTime(config, 1)
|
||||
assert.NotEmpty(t, estimate)
|
||||
assert.True(t, strings.HasSuffix(estimate, "ms") ||
|
||||
strings.HasSuffix(estimate, "s") ||
|
||||
strings.HasSuffix(estimate, "min"))
|
||||
|
||||
// Test different scales
|
||||
estimate = EstimateTime(config, 100)
|
||||
assert.NotEmpty(t, estimate)
|
||||
|
||||
// High security config should take longer
|
||||
highConfig := HighSecurityConfig()
|
||||
highEstimate := EstimateTime(highConfig, 1)
|
||||
assert.NotEmpty(t, highEstimate)
|
||||
}
|
||||
|
||||
func TestDecodeHash(t *testing.T) {
|
||||
validHash := "$argon2id$v=19$m=65536,t=1,p=4$c2FsdA$aGFzaA"
|
||||
|
||||
config, salt, hash, err := decodeHash(validHash)
|
||||
require.NoError(t, err)
|
||||
|
||||
assert.Equal(t, uint32(65536), config.Memory)
|
||||
assert.Equal(t, uint32(1), config.Time)
|
||||
assert.Equal(t, uint8(4), config.Parallelism)
|
||||
assert.Equal(t, []byte("salt"), salt)
|
||||
assert.Equal(t, []byte("hash"), hash)
|
||||
|
||||
// Test invalid formats
|
||||
invalidHashes := []string{
|
||||
"invalid",
|
||||
"$bcrypt$v=19$m=65536,t=1,p=4$salt$hash",
|
||||
"$argon2id$v=18$m=65536,t=1,p=4$salt$hash", // wrong version
|
||||
"$argon2id$v=19$invalid$salt$hash",
|
||||
"$argon2id$v=19$m=65536,t=1,p=4$!invalid!$hash",
|
||||
}
|
||||
|
||||
for _, h := range invalidHashes {
|
||||
_, _, _, err := decodeHash(h)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkDeriveKey(b *testing.B) {
|
||||
configs := map[string]*Config{
|
||||
"light": LightConfig(),
|
||||
"default": DefaultConfig(),
|
||||
"high": HighSecurityConfig(),
|
||||
}
|
||||
|
||||
password := []byte("benchmark-password")
|
||||
salt := []byte("benchmark-salt-at-least-16-bytes")
|
||||
|
||||
for name, config := range configs {
|
||||
b.Run(name, func(b *testing.B) {
|
||||
kdf := New(config)
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_ = kdf.DeriveKey(password, salt)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkHashPassword(b *testing.B) {
|
||||
kdf := New(LightConfig()) // Use light config for benchmarks
|
||||
password := []byte("benchmark-password")
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_, _ = kdf.HashPassword(password)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkVerifyPassword(b *testing.B) {
|
||||
kdf := New(LightConfig())
|
||||
password := []byte("benchmark-password")
|
||||
hash, _ := kdf.HashPassword(password)
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_, _ = VerifyPassword(password, hash)
|
||||
}
|
||||
}
|
||||
|
||||
func TestConcurrentDerivation(t *testing.T) {
|
||||
kdf := New(DefaultConfig())
|
||||
password := []byte("concurrent-test")
|
||||
|
||||
// Generate multiple salts
|
||||
salts := make([][]byte, 10)
|
||||
for i := range salts {
|
||||
salt, err := kdf.GenerateSalt()
|
||||
require.NoError(t, err)
|
||||
salts[i] = salt
|
||||
}
|
||||
|
||||
// Derive keys concurrently
|
||||
results := make([][]byte, len(salts))
|
||||
done := make(chan int, len(salts))
|
||||
|
||||
for i, salt := range salts {
|
||||
go func(idx int, s []byte) {
|
||||
results[idx] = kdf.DeriveKey(password, s)
|
||||
done <- idx
|
||||
}(i, salt)
|
||||
}
|
||||
|
||||
// Wait for all goroutines
|
||||
for i := 0; i < len(salts); i++ {
|
||||
<-done
|
||||
}
|
||||
|
||||
// Verify all keys are different (different salts)
|
||||
for i := 0; i < len(results)-1; i++ {
|
||||
for j := i + 1; j < len(results); j++ {
|
||||
assert.False(t, bytes.Equal(results[i], results[j]))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPerformanceBenchmark(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping performance benchmark in short mode")
|
||||
}
|
||||
|
||||
configs := []struct {
|
||||
name string
|
||||
config *Config
|
||||
maxMs int64
|
||||
}{
|
||||
{"light", LightConfig(), 100},
|
||||
{"default", DefaultConfig(), 500},
|
||||
}
|
||||
|
||||
password := []byte("perf-test")
|
||||
|
||||
for _, tc := range configs {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
kdf := New(tc.config)
|
||||
salt, _ := kdf.GenerateSalt()
|
||||
|
||||
start := time.Now()
|
||||
_ = kdf.DeriveKey(password, salt)
|
||||
elapsed := time.Since(start).Milliseconds()
|
||||
|
||||
t.Logf("%s config took %dms", tc.name, elapsed)
|
||||
assert.Less(t, elapsed, tc.maxMs,
|
||||
"derivation took too long: %dms > %dms", elapsed, tc.maxMs)
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user