feat(auth): implement user and device authentication
This commit is contained in:
@@ -0,0 +1,252 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"cmroubao/backend-api/internal/config"
|
||||
"cmroubao/backend-api/internal/platform/database"
|
||||
"cmroubao/backend-api/internal/platform/migration"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func TestRunCreatesUserWithoutExposingPassword(t *testing.T) {
|
||||
databasePath := migratedDatabase(t)
|
||||
const secret = "correct-horse-password"
|
||||
lookup := testLookup(databasePath, secret)
|
||||
var output bytes.Buffer
|
||||
|
||||
err := run(
|
||||
[]string{"create-user", "ADMIN", "Admin01"},
|
||||
lookup,
|
||||
&output,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("run(create-user) error = %v", err)
|
||||
}
|
||||
if strings.Contains(output.String(), secret) ||
|
||||
!strings.Contains(output.String(), "username=admin01 role=ADMIN") {
|
||||
t.Fatalf("output = %q", output.String())
|
||||
}
|
||||
db, err := database.Open(context.Background(), databasePath)
|
||||
if err != nil {
|
||||
t.Fatalf("database.Open() error = %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
var hash string
|
||||
if err := db.QueryRow(
|
||||
"SELECT password_hash FROM users WHERE username = 'admin01'",
|
||||
).Scan(&hash); err != nil {
|
||||
t.Fatalf("query user: %v", err)
|
||||
}
|
||||
if hash == secret ||
|
||||
bcrypt.CompareHashAndPassword([]byte(hash), []byte(secret)) != nil {
|
||||
t.Fatal("stored password is missing, plaintext, or invalid")
|
||||
}
|
||||
|
||||
output.Reset()
|
||||
err = run(
|
||||
[]string{"create-user", "ADMIN", "admin01"},
|
||||
lookup,
|
||||
&output,
|
||||
)
|
||||
if err == nil ||
|
||||
strings.Contains(err.Error(), secret) ||
|
||||
output.Len() != 0 {
|
||||
t.Fatalf("duplicate error/output = %v / %q", err, output.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunCreatesDeviceAndStoresOnlyTokenHash(t *testing.T) {
|
||||
databasePath := migratedDatabase(t)
|
||||
var output bytes.Buffer
|
||||
|
||||
err := run(
|
||||
[]string{"create-device", "buyer-phone-01"},
|
||||
testLookup(databasePath, ""),
|
||||
&output,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("run(create-device) error = %v", err)
|
||||
}
|
||||
lines := strings.Split(strings.TrimSpace(output.String()), "\n")
|
||||
if len(lines) != 2 ||
|
||||
!strings.HasPrefix(lines[0], "device_id=") ||
|
||||
!strings.HasPrefix(lines[1], "device_token=") {
|
||||
t.Fatalf("output = %q", output.String())
|
||||
}
|
||||
deviceID := strings.TrimPrefix(lines[0], "device_id=")
|
||||
token := strings.TrimPrefix(lines[1], "device_token=")
|
||||
db, err := database.Open(context.Background(), databasePath)
|
||||
if err != nil {
|
||||
t.Fatalf("database.Open() error = %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
var tokenHash string
|
||||
if err := db.QueryRow(
|
||||
"SELECT token_hash FROM devices WHERE id = ?",
|
||||
deviceID,
|
||||
).Scan(&tokenHash); err != nil {
|
||||
t.Fatalf("query device: %v", err)
|
||||
}
|
||||
sum := sha256.Sum256([]byte(token))
|
||||
if tokenHash == token ||
|
||||
tokenHash != hex.EncodeToString(sum[:]) {
|
||||
t.Fatal("stored device token is not the expected hash")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunEnablesAndDisablesUsersAndDevices(t *testing.T) {
|
||||
databasePath := migratedDatabase(t)
|
||||
lookup := testLookup(databasePath, "correct-horse-password")
|
||||
if err := run(
|
||||
[]string{"create-user", "BUYER", "buyer01"},
|
||||
lookup,
|
||||
&bytes.Buffer{},
|
||||
); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
var deviceOutput bytes.Buffer
|
||||
if err := run(
|
||||
[]string{"create-device", "buyer-phone-01"},
|
||||
lookup,
|
||||
&deviceOutput,
|
||||
); err != nil {
|
||||
t.Fatalf("create device: %v", err)
|
||||
}
|
||||
deviceID := strings.TrimPrefix(
|
||||
strings.Split(strings.TrimSpace(deviceOutput.String()), "\n")[0],
|
||||
"device_id=",
|
||||
)
|
||||
|
||||
var output bytes.Buffer
|
||||
if err := run(
|
||||
[]string{"disable-user", "BUYER01"},
|
||||
lookup,
|
||||
&output,
|
||||
); err != nil {
|
||||
t.Fatalf("disable user: %v", err)
|
||||
}
|
||||
if !strings.Contains(output.String(), "username=buyer01 active=false") {
|
||||
t.Fatalf("disable user output = %q", output.String())
|
||||
}
|
||||
output.Reset()
|
||||
if err := run(
|
||||
[]string{"disable-device", deviceID},
|
||||
lookup,
|
||||
&output,
|
||||
); err != nil {
|
||||
t.Fatalf("disable device: %v", err)
|
||||
}
|
||||
|
||||
db, err := database.Open(context.Background(), databasePath)
|
||||
if err != nil {
|
||||
t.Fatalf("database.Open() error = %v", err)
|
||||
}
|
||||
defer db.Close()
|
||||
var userActive, deviceEnabled bool
|
||||
if err := db.QueryRow(
|
||||
"SELECT is_active FROM users WHERE username = 'buyer01'",
|
||||
).Scan(&userActive); err != nil {
|
||||
t.Fatalf("query user status: %v", err)
|
||||
}
|
||||
if err := db.QueryRow(
|
||||
"SELECT is_enabled FROM devices WHERE id = ?",
|
||||
deviceID,
|
||||
).Scan(&deviceEnabled); err != nil {
|
||||
t.Fatalf("query device status: %v", err)
|
||||
}
|
||||
if userActive || deviceEnabled {
|
||||
t.Fatalf(
|
||||
"user/device enabled = %v/%v",
|
||||
userActive,
|
||||
deviceEnabled,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunRejectsWeakPasswordAndInvalidCommands(t *testing.T) {
|
||||
databasePath := migratedDatabase(t)
|
||||
err := run(
|
||||
[]string{"create-user", "BUYER", "buyer"},
|
||||
testLookup(databasePath, "too-short"),
|
||||
&bytes.Buffer{},
|
||||
)
|
||||
if err == nil || err.Error() != "authentication input is invalid" {
|
||||
t.Fatalf("weak password error = %v", err)
|
||||
}
|
||||
for _, arguments := range [][]string{
|
||||
nil,
|
||||
{"create-user", "OWNER", "user"},
|
||||
{"create-device"},
|
||||
{"unknown"},
|
||||
} {
|
||||
if err := run(
|
||||
arguments,
|
||||
testLookup(databasePath, ""),
|
||||
&bytes.Buffer{},
|
||||
); err == nil {
|
||||
t.Fatalf("run(%v) error = nil", arguments)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunRequiresCurrentMigrations(t *testing.T) {
|
||||
databasePath := filepath.Join(t.TempDir(), "pending.db")
|
||||
err := run(
|
||||
[]string{"create-device", "phone"},
|
||||
testLookup(databasePath, ""),
|
||||
&bytes.Buffer{},
|
||||
)
|
||||
if err == nil ||
|
||||
err.Error() != "database migrations are pending; run migrate up" {
|
||||
t.Fatalf("pending migration error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func migratedDatabase(t *testing.T) string {
|
||||
t.Helper()
|
||||
databasePath := filepath.Join(t.TempDir(), "authctl.db")
|
||||
db, err := database.Open(context.Background(), databasePath)
|
||||
if err != nil {
|
||||
t.Fatalf("database.Open() error = %v", err)
|
||||
}
|
||||
runner, err := migration.New(db)
|
||||
if err != nil {
|
||||
t.Fatalf("migration.New() error = %v", err)
|
||||
}
|
||||
if _, err := runner.Up(context.Background()); err != nil {
|
||||
t.Fatalf("migration.Up() error = %v", err)
|
||||
}
|
||||
if err := db.Close(); err != nil {
|
||||
t.Fatalf("db.Close() error = %v", err)
|
||||
}
|
||||
return databasePath
|
||||
}
|
||||
|
||||
func testLookup(
|
||||
databasePath string,
|
||||
password string,
|
||||
) config.LookupEnvironment {
|
||||
return func(name string) (string, bool) {
|
||||
switch name {
|
||||
case config.DatabasePathEnvironment:
|
||||
return databasePath, true
|
||||
case passwordEnvironment:
|
||||
if password == "" {
|
||||
return "", false
|
||||
}
|
||||
return password, true
|
||||
default:
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user