first commit
This commit is contained in:
125
go-api/internal/auth/credentials.go
Normal file
125
go-api/internal/auth/credentials.go
Normal file
@@ -0,0 +1,125 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// ErrInvalidCredentials is the single answer to every failed sign-in.
|
||||
//
|
||||
// It is returned when the email is unknown, when the password is wrong, when
|
||||
// the account has no password set, and when the account is suspended. The
|
||||
// caller cannot tell those apart from the error, which is the point: an error
|
||||
// that distinguishes them is an account-enumeration oracle, and one that says
|
||||
// "this account is suspended" confirms the address is real.
|
||||
//
|
||||
// The distinction is still made — it goes to the server log, via the reason
|
||||
// returned alongside this error.
|
||||
var ErrInvalidCredentials = errors.New("auth: invalid credentials")
|
||||
|
||||
// Reason is why a sign-in failed. For the log, never for the client.
|
||||
type Reason string
|
||||
|
||||
const (
|
||||
ReasonOK Reason = "ok"
|
||||
ReasonNoSuchUser Reason = "no_such_user"
|
||||
ReasonNoPassword Reason = "no_password_set"
|
||||
ReasonNotActive Reason = "user_not_active"
|
||||
ReasonBadPassword Reason = "wrong_password"
|
||||
)
|
||||
|
||||
// Credentials verifies a password against a stored user.
|
||||
type Credentials struct {
|
||||
users UserStore
|
||||
}
|
||||
|
||||
// NewCredentials builds the verifier.
|
||||
func NewCredentials(users UserStore) *Credentials { return &Credentials{users: users} }
|
||||
|
||||
// decoyHash is a real argon2id hash of a value nobody knows.
|
||||
//
|
||||
// It exists to close a timing side channel. Without it, an unknown email
|
||||
// returns as fast as the database can say "no rows" — a millisecond or two —
|
||||
// while a known email spends the ~100ms that argon2id costs by design. That
|
||||
// difference is trivially measurable over a network and turns the login
|
||||
// endpoint into an account-enumeration oracle no matter how carefully the
|
||||
// error messages are worded.
|
||||
//
|
||||
// So every failure that skips the real password check pays for a decoy one
|
||||
// instead. The comparison always fails; the cost is the entire purpose.
|
||||
//
|
||||
// Built once, lazily: it costs a full argon2id derivation, which is worth
|
||||
// paying on the first failed login rather than on every process start.
|
||||
var decoyHash = sync.OnceValue(func() string {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
// A hash of a fixed string is still a fine decoy — its only job is to
|
||||
// take the right amount of time, and it is never compared against
|
||||
// anything a caller supplies.
|
||||
token = "decoy-password-that-is-never-correct"
|
||||
}
|
||||
hash, err := HashPassword(token)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return hash
|
||||
})
|
||||
|
||||
// burnTime performs a password verification that is guaranteed to fail, so a
|
||||
// rejected sign-in costs the same as an accepted one.
|
||||
func burnTime(password string) {
|
||||
if h := decoyHash(); h != "" {
|
||||
_, _ = VerifyPassword(h, password)
|
||||
}
|
||||
}
|
||||
|
||||
// Verify resolves an email and password to a user.
|
||||
//
|
||||
// On success it returns the user and ReasonOK. On any failure it returns
|
||||
// ErrInvalidCredentials, a zero user, and the reason — which the caller should
|
||||
// log and must not send to the client.
|
||||
//
|
||||
// A non-nil error that is NOT ErrInvalidCredentials is an operational failure
|
||||
// (the database is down, a stored hash is corrupt) and should become a 500
|
||||
// rather than a 401: the caller's credentials were never actually judged.
|
||||
func (c *Credentials) Verify(ctx context.Context, email, password string) (User, Reason, error) {
|
||||
user, err := c.users.FindByEmail(ctx, email)
|
||||
if errors.Is(err, ErrUserNotFound) {
|
||||
burnTime(password)
|
||||
return User{}, ReasonNoSuchUser, ErrInvalidCredentials
|
||||
}
|
||||
if err != nil {
|
||||
return User{}, ReasonNoSuchUser, err
|
||||
}
|
||||
|
||||
if user.PasswordHash == "" {
|
||||
// The seeded user is in this state until `setpassword` is run against
|
||||
// it. Refused exactly like a wrong password, at the same cost.
|
||||
burnTime(password)
|
||||
return User{}, ReasonNoPassword, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
ok, err := VerifyPassword(user.PasswordHash, password)
|
||||
if err != nil {
|
||||
// The stored hash could not be read. That is this server's problem,
|
||||
// not the caller's, and must not be reported as a failed login.
|
||||
return User{}, ReasonBadPassword, err
|
||||
}
|
||||
if !ok {
|
||||
return User{}, ReasonBadPassword, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
// Status is checked AFTER the password, and reported the same way.
|
||||
//
|
||||
// Order matters: checking it first would let anyone learn that an address
|
||||
// belongs to a suspended account without knowing its password, because the
|
||||
// refusal would arrive without paying the argon2 cost. Checking it after
|
||||
// means a suspended account is indistinguishable from a wrong password —
|
||||
// same answer, same timing.
|
||||
if !user.IsActive() {
|
||||
return User{}, ReasonNotActive, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
return user, ReasonOK, nil
|
||||
}
|
||||
230
go-api/internal/auth/password.go
Normal file
230
go-api/internal/auth/password.go
Normal file
@@ -0,0 +1,230 @@
|
||||
// Package auth is the authentication foundation: password hashing, session
|
||||
// tokens, and the server-side session lifecycle.
|
||||
//
|
||||
// It deliberately knows nothing about HTTP. There is no handler, no cookie and
|
||||
// no middleware here — those arrive in a later phase and will be written in
|
||||
// terms of this package, not inside it. What lives here is the part that must
|
||||
// be correct regardless of transport: how a password becomes a hash, how a
|
||||
// session token is generated and stored, and when a session stops being valid.
|
||||
//
|
||||
// Two rules hold throughout, and every function below is written to keep them:
|
||||
//
|
||||
// - A raw session token exists in exactly two places: the response that
|
||||
// created it, and the client's cookie. The database holds SHA-256 of it.
|
||||
// - Neither a password nor a token nor a password hash is ever returned in an
|
||||
// error, formatted into a string, or logged. Nothing in this package logs.
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"golang.org/x/crypto/argon2"
|
||||
)
|
||||
|
||||
// Password policy errors. They describe the rule that was broken and never
|
||||
// echo the password back.
|
||||
var (
|
||||
ErrEmptyPassword = errors.New("auth: password is empty")
|
||||
ErrPasswordTooShort = errors.New("auth: password is shorter than the minimum length")
|
||||
ErrPasswordTooLong = errors.New("auth: password is longer than the maximum length")
|
||||
|
||||
// ErrInvalidHash means the stored value is not a hash this package wrote:
|
||||
// wrong prefix, wrong field count, or unparseable parameters.
|
||||
ErrInvalidHash = errors.New("auth: password hash is malformed")
|
||||
|
||||
// ErrIncompatibleVersion means the hash was produced by a future argon2
|
||||
// version this binary cannot verify. Distinguished from ErrInvalidHash
|
||||
// because it is an upgrade problem, not corruption.
|
||||
ErrIncompatibleVersion = errors.New("auth: password hash uses an unsupported argon2 version")
|
||||
)
|
||||
|
||||
const (
|
||||
// MinPasswordLength is measured in bytes, not runes. A byte floor is the
|
||||
// honest one: it is what the KDF consumes, and counting runes would let a
|
||||
// short ASCII password through by way of a generous rune count.
|
||||
MinPasswordLength = 12
|
||||
|
||||
// MaxPasswordLength caps the input. Argon2 has no internal length limit —
|
||||
// unlike bcrypt, it does not silently truncate — so the only reason for a
|
||||
// ceiling is to stop an unbounded body from being hashed at 64 MiB of
|
||||
// memory per attempt. 1 KiB is far above any real passphrase.
|
||||
MaxPasswordLength = 1024
|
||||
)
|
||||
|
||||
// PasswordParams are the argon2id cost parameters.
|
||||
//
|
||||
// They are stored inside every hash this package writes, so a future increase
|
||||
// does not invalidate existing hashes: verification reads the parameters out of
|
||||
// the stored string rather than assuming today's defaults.
|
||||
type PasswordParams struct {
|
||||
// Memory is the KiB of memory the KDF fills. This is the parameter that
|
||||
// makes GPU and ASIC attacks expensive, and the one worth raising first.
|
||||
Memory uint32
|
||||
// Time is the number of passes over that memory.
|
||||
Time uint32
|
||||
// Threads is the parallelism (argon2's `p`).
|
||||
Threads uint8
|
||||
// SaltLength and KeyLength are in bytes.
|
||||
SaltLength uint32
|
||||
KeyLength uint32
|
||||
}
|
||||
|
||||
// DefaultPasswordParams follows the OWASP Password Storage Cheat Sheet's
|
||||
// argon2id recommendation: 64 MiB of memory, 3 iterations, 4 lanes (m=65536,
|
||||
// t=3, p=4). A 16-byte salt and a 32-byte key are the RFC 9106 defaults.
|
||||
//
|
||||
// This costs roughly a tenth of a second per login on developer hardware,
|
||||
// which is the point: it is a cost an attacker pays per guess.
|
||||
var DefaultPasswordParams = PasswordParams{
|
||||
Memory: 64 * 1024,
|
||||
Time: 3,
|
||||
Threads: 4,
|
||||
SaltLength: 16,
|
||||
KeyLength: 32,
|
||||
}
|
||||
|
||||
// HashPassword hashes a plaintext password with the default parameters.
|
||||
//
|
||||
// The returned string is a complete, self-describing PHC record — algorithm,
|
||||
// version, parameters, salt and digest — and is what belongs in
|
||||
// users.password_hash. It is safe to store and unsafe to log.
|
||||
func HashPassword(plain string) (string, error) {
|
||||
return HashPasswordWithParams(plain, DefaultPasswordParams)
|
||||
}
|
||||
|
||||
// HashPasswordWithParams is HashPassword with explicit cost parameters. Tests
|
||||
// use it to run at a cost that does not dominate the test suite; production
|
||||
// code should call HashPassword.
|
||||
func HashPasswordWithParams(plain string, p PasswordParams) (string, error) {
|
||||
if err := ValidatePassword(plain); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if p.SaltLength == 0 || p.KeyLength == 0 || p.Memory == 0 || p.Time == 0 || p.Threads == 0 {
|
||||
return "", fmt.Errorf("auth: argon2id parameters must all be non-zero")
|
||||
}
|
||||
|
||||
salt := make([]byte, p.SaltLength)
|
||||
if _, err := rand.Read(salt); err != nil {
|
||||
// crypto/rand failing is not recoverable and must never fall back to a
|
||||
// weaker source: a predictable salt defeats the whole construction.
|
||||
return "", fmt.Errorf("auth: read salt: %w", err)
|
||||
}
|
||||
|
||||
key := argon2.IDKey([]byte(plain), salt, p.Time, p.Memory, p.Threads, p.KeyLength)
|
||||
|
||||
// The PHC string format, as produced by the reference implementation:
|
||||
// $argon2id$v=19$m=65536,t=3,p=4$<b64 salt>$<b64 key>
|
||||
// Standard base64 without padding, which is what the format specifies.
|
||||
return fmt.Sprintf("$argon2id$v=%d$m=%d,t=%d,p=%d$%s$%s",
|
||||
argon2.Version, p.Memory, p.Time, p.Threads,
|
||||
base64.RawStdEncoding.EncodeToString(salt),
|
||||
base64.RawStdEncoding.EncodeToString(key),
|
||||
), nil
|
||||
}
|
||||
|
||||
// VerifyPassword reports whether plain is the password behind encoded.
|
||||
//
|
||||
// A false return with a nil error is the ordinary "wrong password" answer. A
|
||||
// non-nil error means the *stored hash* could not be read, which is an
|
||||
// operational problem rather than a failed login, and callers should tell the
|
||||
// two apart: the first is a 401, the second is a 500.
|
||||
//
|
||||
// The digest comparison is constant-time. The length and parameter checks
|
||||
// before it are not, and do not need to be: they depend only on the stored
|
||||
// hash, never on the supplied password.
|
||||
func VerifyPassword(encoded, plain string) (bool, error) {
|
||||
p, salt, want, err := DecodePasswordHash(encoded)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
// No policy check on `plain` here. A password that predates a tightened
|
||||
// minimum length must still be able to log in; the policy applies when a
|
||||
// password is set, which is where ValidatePassword is called.
|
||||
if len(plain) > MaxPasswordLength {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
got := argon2.IDKey([]byte(plain), salt, p.Time, p.Memory, p.Threads, p.KeyLength)
|
||||
return subtle.ConstantTimeCompare(got, want) == 1, nil
|
||||
}
|
||||
|
||||
// ValidatePassword applies the policy for setting a new password.
|
||||
func ValidatePassword(plain string) error {
|
||||
switch {
|
||||
case len(plain) == 0:
|
||||
return ErrEmptyPassword
|
||||
case len(plain) < MinPasswordLength:
|
||||
return ErrPasswordTooShort
|
||||
case len(plain) > MaxPasswordLength:
|
||||
return ErrPasswordTooLong
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DecodePasswordHash parses a PHC argon2id record back into its parts.
|
||||
//
|
||||
// Exported so that a future re-hash-on-login path can ask whether a stored hash
|
||||
// was written with weaker parameters than today's default and upgrade it. It
|
||||
// returns the salt and digest, never the password.
|
||||
func DecodePasswordHash(encoded string) (p PasswordParams, salt, key []byte, err error) {
|
||||
// $argon2id$v=19$m=65536,t=3,p=4$<salt>$<key> splits into six fields, the
|
||||
// first of which is empty because the string starts with the separator.
|
||||
parts := strings.Split(encoded, "$")
|
||||
if len(parts) != 6 || parts[0] != "" {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if parts[1] != "argon2id" {
|
||||
// bcrypt, argon2i and argon2d all land here. This package writes and
|
||||
// reads argon2id and nothing else; a different algorithm is a
|
||||
// migration decision, not something to guess at during a login.
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
|
||||
var version int
|
||||
if _, err := fmt.Sscanf(parts[2], "v=%d", &version); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if version != argon2.Version {
|
||||
return p, nil, nil, ErrIncompatibleVersion
|
||||
}
|
||||
|
||||
if _, err := fmt.Sscanf(parts[3], "m=%d,t=%d,p=%d", &p.Memory, &p.Time, &p.Threads); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if p.Memory == 0 || p.Time == 0 || p.Threads == 0 {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
|
||||
if salt, err = base64.RawStdEncoding.DecodeString(parts[4]); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if key, err = base64.RawStdEncoding.DecodeString(parts[5]); err != nil {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
if len(salt) == 0 || len(key) == 0 {
|
||||
return p, nil, nil, ErrInvalidHash
|
||||
}
|
||||
|
||||
p.SaltLength = uint32(len(salt))
|
||||
p.KeyLength = uint32(len(key))
|
||||
return p, salt, key, nil
|
||||
}
|
||||
|
||||
// NeedsRehash reports whether a stored hash was written with parameters weaker
|
||||
// than want, so a successful login can transparently upgrade it.
|
||||
//
|
||||
// Unused in Phase 3B — there is no login yet — and exported now because the
|
||||
// judgement belongs beside the format that encodes the parameters.
|
||||
func NeedsRehash(encoded string, want PasswordParams) bool {
|
||||
p, _, _, err := DecodePasswordHash(encoded)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
return p.Memory < want.Memory || p.Time < want.Time ||
|
||||
p.KeyLength < want.KeyLength || p.SaltLength < want.SaltLength
|
||||
}
|
||||
254
go-api/internal/auth/password_test.go
Normal file
254
go-api/internal/auth/password_test.go
Normal file
@@ -0,0 +1,254 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// testParams runs argon2id at a cost that is still real but does not make the
|
||||
// suite crawl. Every property under test — salting, verification, the encoded
|
||||
// format — is independent of the cost, and DefaultPasswordParams is asserted
|
||||
// separately in TestDefaultPasswordParamsMeetOWASP.
|
||||
var testParams = PasswordParams{Memory: 8 * 1024, Time: 1, Threads: 2, SaltLength: 16, KeyLength: 32}
|
||||
|
||||
const goodPassword = "correct-horse-battery-staple"
|
||||
|
||||
// 1. Hash generation produces a well-formed, self-describing argon2id record.
|
||||
func TestHashPasswordProducesArgon2idRecord(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("HashPasswordWithParams: %v", err)
|
||||
}
|
||||
|
||||
if !strings.HasPrefix(hash, "$argon2id$") {
|
||||
t.Fatalf("hash is not argon2id: %q", firstField(hash))
|
||||
}
|
||||
// bcrypt would be "$2a$"/"$2b$"; argon2i and argon2d are different KDFs.
|
||||
// The decision was argon2id specifically, so assert it rather than merely
|
||||
// "some hash was produced".
|
||||
if parts := strings.Split(hash, "$"); len(parts) != 6 {
|
||||
t.Fatalf("hash has %d fields, want 6 (PHC format)", len(parts))
|
||||
}
|
||||
|
||||
got, salt, key, err := DecodePasswordHash(hash)
|
||||
if err != nil {
|
||||
t.Fatalf("DecodePasswordHash: %v", err)
|
||||
}
|
||||
if got.Memory != testParams.Memory || got.Time != testParams.Time || got.Threads != testParams.Threads {
|
||||
t.Errorf("decoded params = m=%d,t=%d,p=%d, want m=%d,t=%d,p=%d",
|
||||
got.Memory, got.Time, got.Threads, testParams.Memory, testParams.Time, testParams.Threads)
|
||||
}
|
||||
if len(salt) != int(testParams.SaltLength) {
|
||||
t.Errorf("salt is %d bytes, want %d", len(salt), testParams.SaltLength)
|
||||
}
|
||||
if len(key) != int(testParams.KeyLength) {
|
||||
t.Errorf("key is %d bytes, want %d", len(key), testParams.KeyLength)
|
||||
}
|
||||
// The whole point of the format: the hash must not contain the password.
|
||||
if strings.Contains(hash, goodPassword) {
|
||||
t.Error("the encoded hash contains the plaintext password")
|
||||
}
|
||||
}
|
||||
|
||||
// 2. The correct password verifies.
|
||||
func TestVerifyPasswordAcceptsTheCorrectPassword(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
ok, err := VerifyPassword(hash, goodPassword)
|
||||
if err != nil {
|
||||
t.Fatalf("VerifyPassword: %v", err)
|
||||
}
|
||||
if !ok {
|
||||
t.Fatal("the correct password did not verify")
|
||||
}
|
||||
}
|
||||
|
||||
// 3. An incorrect password is rejected — including the near misses that a
|
||||
// sloppy comparison would let through.
|
||||
func TestVerifyPasswordRejectsIncorrectPasswords(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
|
||||
wrong := map[string]string{
|
||||
"different": "incorrect-horse-battery-staple",
|
||||
"empty": "",
|
||||
"prefix": goodPassword[:len(goodPassword)-1],
|
||||
"suffix appended": goodPassword + "x",
|
||||
"case flipped": strings.ToUpper(goodPassword),
|
||||
"whitespace": " " + goodPassword,
|
||||
}
|
||||
for name, candidate := range wrong {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
ok, err := VerifyPassword(hash, candidate)
|
||||
if err != nil {
|
||||
t.Fatalf("VerifyPassword returned an error for a wrong password: %v", err)
|
||||
}
|
||||
if ok {
|
||||
t.Error("a wrong password verified")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 4. Different passwords produce different hashes — and so does the SAME
|
||||
// password hashed twice, which is the stronger property and the one that
|
||||
// actually depends on the salt being random.
|
||||
func TestHashPasswordIsSaltedPerCall(t *testing.T) {
|
||||
a, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash a: %v", err)
|
||||
}
|
||||
b, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash b: %v", err)
|
||||
}
|
||||
if a == b {
|
||||
t.Fatal("hashing the same password twice produced identical hashes; the salt is not random")
|
||||
}
|
||||
// Both must still verify: a per-call salt is only useful if it travels
|
||||
// with the hash.
|
||||
for i, h := range []string{a, b} {
|
||||
ok, err := VerifyPassword(h, goodPassword)
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("hash %d did not verify its own password (ok=%v err=%v)", i, ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
c, err := HashPasswordWithParams("a-completely-different-password", testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash c: %v", err)
|
||||
}
|
||||
if c == a {
|
||||
t.Error("different passwords produced identical hashes")
|
||||
}
|
||||
// And a hash must not verify a password it was not made from.
|
||||
if ok, _ := VerifyPassword(c, goodPassword); ok {
|
||||
t.Error("a hash verified a password it was not derived from")
|
||||
}
|
||||
}
|
||||
|
||||
// 5a. Empty and out-of-policy passwords are refused at hashing time.
|
||||
func TestHashPasswordRejectsInvalidPasswords(t *testing.T) {
|
||||
cases := map[string]struct {
|
||||
password string
|
||||
want error
|
||||
}{
|
||||
"empty": {"", ErrEmptyPassword},
|
||||
"too short": {strings.Repeat("a", MinPasswordLength-1), ErrPasswordTooShort},
|
||||
"too long": {strings.Repeat("a", MaxPasswordLength+1), ErrPasswordTooLong},
|
||||
}
|
||||
for name, tc := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
hash, err := HashPasswordWithParams(tc.password, testParams)
|
||||
if err != tc.want {
|
||||
t.Fatalf("error = %v, want %v", err, tc.want)
|
||||
}
|
||||
if hash != "" {
|
||||
t.Error("a hash was returned alongside the error")
|
||||
}
|
||||
// The rejection must not quote the input back.
|
||||
if tc.password != "" && err != nil && strings.Contains(err.Error(), tc.password) {
|
||||
t.Error("the error message contains the password")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// The boundary itself is allowed: the rule is "at least MinPasswordLength".
|
||||
if _, err := HashPasswordWithParams(strings.Repeat("a", MinPasswordLength), testParams); err != nil {
|
||||
t.Errorf("a password of exactly the minimum length was rejected: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 5b. A malformed *stored* hash is an error, not a silent "wrong password".
|
||||
// The distinction matters: one is a 401, the other is a 500.
|
||||
func TestVerifyPasswordRejectsMalformedHashes(t *testing.T) {
|
||||
valid, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
fields := strings.Split(valid, "$")
|
||||
|
||||
bad := map[string]string{
|
||||
"empty": "",
|
||||
"not a phc string": "not-a-hash",
|
||||
"bcrypt": "$2a$10$N9qo8uLOickgx2ZMRZoMyeIjZAgcfl7p92ldGxad68LJZdL17lhWy",
|
||||
"argon2i": strings.Replace(valid, "argon2id", "argon2i", 1),
|
||||
"too few fields": strings.Join(fields[:5], "$"),
|
||||
"unparseable params": strings.Replace(valid, fields[3], "m=x,t=y,p=z", 1),
|
||||
"zero memory": strings.Replace(valid, fields[3], "m=0,t=1,p=2", 1),
|
||||
"bad version": strings.Replace(valid, fields[2], "v=notanumber", 1),
|
||||
"salt not base64": strings.Replace(valid, fields[4], "!!!!not base64!!!!", 1),
|
||||
}
|
||||
for name, encoded := range bad {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
ok, err := VerifyPassword(encoded, goodPassword)
|
||||
if err == nil {
|
||||
t.Fatal("a malformed hash verified without an error")
|
||||
}
|
||||
if ok {
|
||||
t.Error("a malformed hash reported a successful verification")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// A future argon2 version is reported as its own error, because it is an
|
||||
// upgrade problem rather than corruption.
|
||||
future := strings.Replace(valid, fields[2], "v=99", 1)
|
||||
if _, err := VerifyPassword(future, goodPassword); err != ErrIncompatibleVersion {
|
||||
t.Errorf("error for a future version = %v, want ErrIncompatibleVersion", err)
|
||||
}
|
||||
}
|
||||
|
||||
// The defaults are a security decision, so they are asserted rather than
|
||||
// assumed: OWASP's argon2id recommendation is m=65536 (64 MiB), t=3, p=4.
|
||||
func TestDefaultPasswordParamsMeetOWASP(t *testing.T) {
|
||||
p := DefaultPasswordParams
|
||||
if p.Memory < 64*1024 {
|
||||
t.Errorf("Memory = %d KiB, want at least 65536", p.Memory)
|
||||
}
|
||||
if p.Time < 3 {
|
||||
t.Errorf("Time = %d, want at least 3", p.Time)
|
||||
}
|
||||
if p.Threads < 1 {
|
||||
t.Errorf("Threads = %d, want at least 1", p.Threads)
|
||||
}
|
||||
if p.SaltLength < 16 {
|
||||
t.Errorf("SaltLength = %d, want at least 16", p.SaltLength)
|
||||
}
|
||||
if p.KeyLength < 32 {
|
||||
t.Errorf("KeyLength = %d, want at least 32", p.KeyLength)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNeedsRehash(t *testing.T) {
|
||||
weak, err := HashPasswordWithParams(goodPassword, testParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
if !NeedsRehash(weak, DefaultPasswordParams) {
|
||||
t.Error("a hash below the default cost was not flagged for rehashing")
|
||||
}
|
||||
strong, err := HashPasswordWithParams(goodPassword, DefaultPasswordParams)
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
if NeedsRehash(strong, DefaultPasswordParams) {
|
||||
t.Error("a hash at the default cost was flagged for rehashing")
|
||||
}
|
||||
if !NeedsRehash("not-a-hash", DefaultPasswordParams) {
|
||||
t.Error("an unreadable hash should be flagged for rehashing")
|
||||
}
|
||||
}
|
||||
|
||||
// firstField is used only to report a failure without dumping a whole hash.
|
||||
func firstField(hash string) string {
|
||||
parts := strings.SplitN(hash, "$", 3)
|
||||
if len(parts) < 2 {
|
||||
return hash
|
||||
}
|
||||
return "$" + parts[1] + "$"
|
||||
}
|
||||
376
go-api/internal/auth/schema_test.go
Normal file
376
go-api/internal/auth/schema_test.go
Normal file
@@ -0,0 +1,376 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sort"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// These tests drive migration 000004 itself: what it creates, that it can be
|
||||
// rolled back, and that rolling it back and re-applying it leaves the schema
|
||||
// where it started. They also cover the one change 000004 makes to an existing
|
||||
// table — the global unique index on users.email.
|
||||
//
|
||||
// Every one runs in its own throwaway database. Nothing here can reach the
|
||||
// development database: testutil builds the name from its own prefix.
|
||||
|
||||
const migration4Up = "000004_auth_sessions.up.sql"
|
||||
const migration4Down = "000004_auth_sessions.down.sql"
|
||||
|
||||
// 14. Email is unique across the whole table, not merely within one
|
||||
// organization. This is what makes "log in with your email" answerable.
|
||||
func TestUsersEmailIsGloballyUnique(t *testing.T) {
|
||||
f := newFixture(t, "schema_email_unique")
|
||||
|
||||
// A second organization. Under 000001's (org_id, email) key alone, the
|
||||
// insert below would have been perfectly legal.
|
||||
var otherOrg string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
"Second Org", "second-org").Scan(&otherOrg); err != nil {
|
||||
t.Fatalf("create the second organization: %v", err)
|
||||
}
|
||||
|
||||
_, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, $3)`,
|
||||
otherOrg, "session-owner@example.test", "Impostor")
|
||||
if err == nil {
|
||||
t.Fatal("the same email was accepted in a second organization; login would be ambiguous")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "users_email_global_key") {
|
||||
t.Errorf("the rejection came from %v, want a users_email_global_key violation", err)
|
||||
}
|
||||
|
||||
// citext makes the index case-insensitive, which is what a login form
|
||||
// needs: nobody should be able to register Demo@… beside demo@….
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, $3)`,
|
||||
otherOrg, "SESSION-OWNER@EXAMPLE.TEST", "Impostor"); err == nil {
|
||||
t.Error("a case-variant of an existing email was accepted")
|
||||
}
|
||||
|
||||
// A genuinely different email in the second organization is still fine —
|
||||
// the index constrains duplicates, not multi-tenancy.
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, $3)`,
|
||||
otherOrg, "someone-else@example.test", "Colleague"); err != nil {
|
||||
t.Errorf("a distinct email in a second organization was rejected: %v", err)
|
||||
}
|
||||
|
||||
// The index is unique and covers exactly users(email).
|
||||
var isUnique bool
|
||||
var definition string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT i.indisunique, pg_get_indexdef(i.indexrelid)
|
||||
FROM pg_index i
|
||||
JOIN pg_class c ON c.oid = i.indexrelid
|
||||
WHERE c.relname = 'users_email_global_key'`).Scan(&isUnique, &definition); err != nil {
|
||||
t.Fatalf("read users_email_global_key: %v", err)
|
||||
}
|
||||
if !isUnique {
|
||||
t.Error("users_email_global_key is not a unique index")
|
||||
}
|
||||
if !strings.Contains(definition, "(email)") {
|
||||
t.Errorf("users_email_global_key covers %q, want (email)", definition)
|
||||
}
|
||||
}
|
||||
|
||||
// 15 and 17. Applying every migration in order produces the sessions table
|
||||
// this package needs, and leaves everything the earlier migrations built.
|
||||
func TestMigrationUpBuildsTheSessionsSchema(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, "schema_up")
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
// Columns and types.
|
||||
want := map[string]string{
|
||||
"id": "uuid",
|
||||
"user_id": "uuid",
|
||||
"token_hash": "text",
|
||||
"expires_at": "timestamp with time zone",
|
||||
"absolute_expires_at": "timestamp with time zone",
|
||||
"created_date": "timestamp with time zone",
|
||||
"last_seen_at": "timestamp with time zone",
|
||||
}
|
||||
rows, err := pool.Query(ctx,
|
||||
`SELECT column_name, data_type, is_nullable
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = 'public' AND table_name = 'sessions'`)
|
||||
if err != nil {
|
||||
t.Fatalf("read the sessions columns: %v", err)
|
||||
}
|
||||
got := map[string]string{}
|
||||
for rows.Next() {
|
||||
var name, dataType, nullable string
|
||||
if err := rows.Scan(&name, &dataType, &nullable); err != nil {
|
||||
t.Fatalf("scan: %v", err)
|
||||
}
|
||||
got[name] = dataType
|
||||
// Every column is required. A nullable expires_at would be a session
|
||||
// with no deadline at all.
|
||||
if nullable != "NO" {
|
||||
t.Errorf("sessions.%s is nullable; every session column is required", name)
|
||||
}
|
||||
}
|
||||
rows.Close()
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatalf("read the sessions columns: %v", err)
|
||||
}
|
||||
if len(got) == 0 {
|
||||
t.Fatal("the sessions table does not exist after the migrations")
|
||||
}
|
||||
for name, wantType := range want {
|
||||
if gotType, ok := got[name]; !ok {
|
||||
t.Errorf("sessions.%s is missing", name)
|
||||
} else if gotType != wantType {
|
||||
t.Errorf("sessions.%s is %s, want %s", name, gotType, wantType)
|
||||
}
|
||||
}
|
||||
for name := range got {
|
||||
if _, expected := want[name]; !expected {
|
||||
t.Errorf("sessions has an unexpected column %q", name)
|
||||
}
|
||||
}
|
||||
|
||||
// Constraints: the primary key, the uniqueness that makes a token identify
|
||||
// one session, and the checks that keep a raw token and an immortal
|
||||
// session out of the table.
|
||||
for _, c := range []struct{ name, kind string }{
|
||||
{"sessions_pkey", "p"},
|
||||
{"sessions_token_hash_key", "u"},
|
||||
{"sessions_token_hash_sha256", "c"},
|
||||
{"sessions_absolute_after_created", "c"},
|
||||
{"sessions_within_absolute", "c"},
|
||||
} {
|
||||
var kind string
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT contype::text FROM pg_constraint
|
||||
WHERE conrelid = 'public.sessions'::regclass AND conname = $1`, c.name).Scan(&kind); err != nil {
|
||||
t.Errorf("constraint %s is missing: %v", c.name, err)
|
||||
continue
|
||||
}
|
||||
if kind != c.kind {
|
||||
t.Errorf("constraint %s is of type %q, want %q", c.name, kind, c.kind)
|
||||
}
|
||||
}
|
||||
|
||||
// The foreign key, and that it cascades. ON DELETE NO ACTION here would
|
||||
// mean a deleted user keeps a working session.
|
||||
var fkTarget, onDelete string
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT confrelid::regclass::text, confdeltype::text
|
||||
FROM pg_constraint
|
||||
WHERE conrelid = 'public.sessions'::regclass AND contype = 'f'`).Scan(&fkTarget, &onDelete); err != nil {
|
||||
t.Fatalf("read the sessions foreign key: %v", err)
|
||||
}
|
||||
if fkTarget != "users" {
|
||||
t.Errorf("the foreign key points at %s, want users", fkTarget)
|
||||
}
|
||||
if onDelete != "c" {
|
||||
t.Errorf("the foreign key is ON DELETE %q, want \"c\" (CASCADE)", onDelete)
|
||||
}
|
||||
|
||||
// Indexes. token_hash is indexed by its UNIQUE constraint — that index is
|
||||
// the lookup path — plus the two the sweep and per-user revocation need.
|
||||
indexes := indexNames(ctx, t, pool, "sessions")
|
||||
for _, name := range []string{"sessions_pkey", "sessions_token_hash_key", "sessions_user_idx", "sessions_expires_idx"} {
|
||||
if !contains(indexes, name) {
|
||||
t.Errorf("index %s is missing; sessions has %v", name, indexes)
|
||||
}
|
||||
}
|
||||
|
||||
// 17. The earlier migrations are untouched: the tables they built are all
|
||||
// still here, and so is the org-scoped uniqueness 000001 declared.
|
||||
//
|
||||
// Named rather than counted. This was a count of 18 — the 17 tables from
|
||||
// 000001 plus sessions — which asserted the right thing in the wrong way:
|
||||
// it broke when 000005 added two unrelated tables, and it would have stayed
|
||||
// green if 000004 had dropped one table and added another. Listing them is
|
||||
// both more precise and stable across later migrations.
|
||||
for _, table := range []string{
|
||||
"organizations", "users", "user_preferences", "role_categories",
|
||||
"certifications", "badges", "courses", "learning_paths", "job_postings",
|
||||
"worker_profiles", "job_applications", "ai_interviews", "staff",
|
||||
"assignments", "shift_records", "evidence", "user_activity",
|
||||
"sessions",
|
||||
} {
|
||||
if !tableExists(ctx, t, pool, table) {
|
||||
t.Errorf("%s is missing after the migrations", table)
|
||||
}
|
||||
}
|
||||
var orgScoped int
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM pg_constraint
|
||||
WHERE conrelid = 'public.users'::regclass AND conname = 'users_org_email_key'`).Scan(&orgScoped); err != nil {
|
||||
t.Fatalf("read users_org_email_key: %v", err)
|
||||
}
|
||||
if orgScoped != 1 {
|
||||
t.Error("000004 removed users_org_email_key; it should leave the existing table alone")
|
||||
}
|
||||
}
|
||||
|
||||
// 16. 000004 rolls back cleanly, and re-applies afterwards. A migration that
|
||||
// cannot be reversed is a migration nobody can safely deploy.
|
||||
func TestMigration000004IsReversible(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, "schema_down")
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
if !tableExists(ctx, t, pool, "sessions") {
|
||||
t.Fatal("sessions does not exist before the rollback")
|
||||
}
|
||||
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, migration4Down); err != nil {
|
||||
t.Fatalf("apply %s: %v", migration4Down, err)
|
||||
}
|
||||
|
||||
if tableExists(ctx, t, pool, "sessions") {
|
||||
t.Error("sessions survived the rollback")
|
||||
}
|
||||
if indexExists(ctx, t, pool, "users_email_global_key") {
|
||||
t.Error("users_email_global_key survived the rollback")
|
||||
}
|
||||
|
||||
// The rollback must touch nothing else. users is still here, still has the
|
||||
// constraint that predates 000004, and still has its rows.
|
||||
if !tableExists(ctx, t, pool, "users") {
|
||||
t.Fatal("the rollback dropped the users table")
|
||||
}
|
||||
var orgScoped int
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM pg_constraint
|
||||
WHERE conrelid = 'public.users'::regclass AND conname = 'users_org_email_key'`).Scan(&orgScoped); err != nil {
|
||||
t.Fatalf("read users_org_email_key: %v", err)
|
||||
}
|
||||
if orgScoped != 1 {
|
||||
t.Error("the rollback removed users_org_email_key, which it did not create")
|
||||
}
|
||||
// With the global index gone, the pre-000004 rule is back in force: the
|
||||
// same email in two organizations is legal again.
|
||||
var orgA, orgB string
|
||||
if err := pool.QueryRow(ctx, `INSERT INTO organizations (name, slug) VALUES ('A','a') RETURNING id::text`).Scan(&orgA); err != nil {
|
||||
t.Fatalf("create org A: %v", err)
|
||||
}
|
||||
if err := pool.QueryRow(ctx, `INSERT INTO organizations (name, slug) VALUES ('B','b') RETURNING id::text`).Scan(&orgB); err != nil {
|
||||
t.Fatalf("create org B: %v", err)
|
||||
}
|
||||
for _, org := range []string{orgA, orgB} {
|
||||
if _, err := pool.Exec(ctx,
|
||||
`INSERT INTO users (org_id, email, full_name) VALUES ($1::uuid, $2::citext, '')`,
|
||||
org, "shared@example.test"); err != nil {
|
||||
t.Fatalf("insert after the rollback: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Re-applying now must fail, because the data violates the uniqueness the
|
||||
// migration is about to declare — and it must fail without half-applying.
|
||||
// That is the honest behaviour: the operator has duplicates to resolve.
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, migration4Up); err == nil {
|
||||
t.Fatal("000004 applied over duplicate emails; the index would not be unique")
|
||||
}
|
||||
if tableExists(ctx, t, pool, "sessions") {
|
||||
t.Error("the failed migration left the sessions table behind; it is not atomic")
|
||||
}
|
||||
|
||||
// Resolve the duplicate and it applies cleanly, restoring exactly what the
|
||||
// rollback removed.
|
||||
if _, err := pool.Exec(ctx, `DELETE FROM users WHERE org_id = $1::uuid`, orgB); err != nil {
|
||||
t.Fatalf("remove the duplicate: %v", err)
|
||||
}
|
||||
if err := testutil.ApplyMigration(ctx, t, pool, migration4Up); err != nil {
|
||||
t.Fatalf("re-apply %s: %v", migration4Up, err)
|
||||
}
|
||||
if !tableExists(ctx, t, pool, "sessions") {
|
||||
t.Error("sessions did not come back")
|
||||
}
|
||||
if !indexExists(ctx, t, pool, "users_email_global_key") {
|
||||
t.Error("users_email_global_key did not come back")
|
||||
}
|
||||
}
|
||||
|
||||
// Every migration has a matching down file, so any of them can be reversed.
|
||||
func TestEveryMigrationHasADownFile(t *testing.T) {
|
||||
ups := testutil.MigrationFiles(t, ".up.sql")
|
||||
downs := testutil.MigrationFiles(t, ".down.sql")
|
||||
if len(ups) != len(downs) {
|
||||
t.Fatalf("%d up migrations and %d down migrations", len(ups), len(downs))
|
||||
}
|
||||
for i, up := range ups {
|
||||
want := strings.TrimSuffix(up, ".up.sql") + ".down.sql"
|
||||
if downs[i] != want {
|
||||
t.Errorf("%s has no matching down migration (found %s)", up, downs[i])
|
||||
}
|
||||
}
|
||||
// 000004 is still present, with its pair. This used to also assert that it
|
||||
// was the NEWEST and that there were exactly four migrations — a snapshot
|
||||
// that this test's own comment predicted would need revisiting, and which
|
||||
// 000005 duly broke. The count belongs to whichever phase added the newest
|
||||
// migration (see TestMigrationPairsIncluding000005), so it is asserted
|
||||
// there and not here. What this phase cares about — that the migration it
|
||||
// added is intact and reversible — is unchanged and still checked.
|
||||
if !contains(ups, migration4Up) {
|
||||
t.Errorf("%s is missing from the migrations directory", migration4Up)
|
||||
}
|
||||
if !contains(downs, migration4Down) {
|
||||
t.Errorf("%s is missing from the migrations directory", migration4Down)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── introspection helpers ──────────────────────────────────────────────── */
|
||||
|
||||
func tableExists(ctx context.Context, t *testing.T, pool *pgxpool.Pool, name string) bool {
|
||||
t.Helper()
|
||||
var reg *string
|
||||
if err := pool.QueryRow(ctx, `SELECT to_regclass('public.' || $1)::text`, name).Scan(®); err != nil {
|
||||
t.Fatalf("to_regclass(%s): %v", name, err)
|
||||
}
|
||||
return reg != nil
|
||||
}
|
||||
|
||||
func indexExists(ctx context.Context, t *testing.T, pool *pgxpool.Pool, name string) bool {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := pool.QueryRow(ctx,
|
||||
`SELECT count(*)::int FROM pg_indexes WHERE schemaname = 'public' AND indexname = $1`,
|
||||
name).Scan(&n); err != nil {
|
||||
t.Fatalf("look up index %s: %v", name, err)
|
||||
}
|
||||
return n > 0
|
||||
}
|
||||
|
||||
func indexNames(ctx context.Context, t *testing.T, pool *pgxpool.Pool, table string) []string {
|
||||
t.Helper()
|
||||
rows, err := pool.Query(ctx,
|
||||
`SELECT indexname FROM pg_indexes WHERE schemaname = 'public' AND tablename = $1`, table)
|
||||
if err != nil {
|
||||
t.Fatalf("list the indexes on %s: %v", table, err)
|
||||
}
|
||||
defer rows.Close()
|
||||
var names []string
|
||||
for rows.Next() {
|
||||
var name string
|
||||
if err := rows.Scan(&name); err != nil {
|
||||
t.Fatalf("scan: %v", err)
|
||||
}
|
||||
names = append(names, name)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
t.Fatalf("list the indexes on %s: %v", table, err)
|
||||
}
|
||||
sort.Strings(names)
|
||||
return names
|
||||
}
|
||||
|
||||
func contains(haystack []string, needle string) bool {
|
||||
for _, s := range haystack {
|
||||
if s == needle {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
324
go-api/internal/auth/session.go
Normal file
324
go-api/internal/auth/session.go
Normal file
@@ -0,0 +1,324 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Session lifecycle errors.
|
||||
var (
|
||||
// ErrSessionNotFound means no row matched the token hash. It is returned
|
||||
// for an unknown token and for a well-formed token that has been revoked;
|
||||
// a caller must not distinguish the two to the client.
|
||||
ErrSessionNotFound = errors.New("auth: session not found")
|
||||
|
||||
// ErrSessionExpired means a row matched but is no longer valid, by either
|
||||
// the sliding or the absolute deadline.
|
||||
ErrSessionExpired = errors.New("auth: session expired")
|
||||
)
|
||||
|
||||
// Session is one row of the sessions table.
|
||||
//
|
||||
// TokenHash is the SHA-256 of the token, never the token. There is no field
|
||||
// here that can hold the raw secret, by design: the only place it exists after
|
||||
// Issue returns is the caller's cookie.
|
||||
type Session struct {
|
||||
ID string
|
||||
UserID string
|
||||
|
||||
// TokenHash is lowercase hex SHA-256. See HashToken.
|
||||
TokenHash string
|
||||
|
||||
// ExpiresAt is the sliding deadline; it moves forward as the session is
|
||||
// used, never past AbsoluteExpiresAt.
|
||||
ExpiresAt time.Time
|
||||
|
||||
// AbsoluteExpiresAt is fixed when the session is created and never moves.
|
||||
AbsoluteExpiresAt time.Time
|
||||
|
||||
CreatedDate time.Time
|
||||
LastSeenAt time.Time
|
||||
}
|
||||
|
||||
// IsExpired reports whether the session is dead at the given instant, by
|
||||
// either deadline. The absolute one is checked as well as the sliding one
|
||||
// precisely because the sliding one can be moved.
|
||||
func (s Session) IsExpired(now time.Time) bool {
|
||||
return !now.Before(s.ExpiresAt) || !now.Before(s.AbsoluteExpiresAt)
|
||||
}
|
||||
|
||||
// Policy is how long a session lives.
|
||||
//
|
||||
// Two pairs of durations, because "Remember Me" is a different risk than a
|
||||
// session on a shared machine, and because a sliding window alone can be slid
|
||||
// forever.
|
||||
type Policy struct {
|
||||
// IdleLifetime is how long a normal session survives without being used.
|
||||
IdleLifetime time.Duration
|
||||
// AbsoluteLifetime caps a normal session's total life regardless of use.
|
||||
AbsoluteLifetime time.Duration
|
||||
|
||||
// RememberIdleLifetime and RememberAbsoluteLifetime are the same two
|
||||
// bounds for a session created with Remember Me.
|
||||
RememberIdleLifetime time.Duration
|
||||
RememberAbsoluteLifetime time.Duration
|
||||
}
|
||||
|
||||
// DefaultPolicy implements the Phase 3B session-lifetime decision.
|
||||
//
|
||||
// normal login 12 hours idle, capped at 24 hours of total life. Twelve hours
|
||||
// covers a working day; the daily cap means an unattended tab
|
||||
// cannot be slid along indefinitely.
|
||||
// Remember Me 30 days idle, capped at 90 days. The user asked to stay
|
||||
// signed in; 90 days is the point at which they re-prove it.
|
||||
//
|
||||
// The idle bound is what expires an abandoned session. The absolute bound is
|
||||
// what guarantees no session lives forever.
|
||||
var DefaultPolicy = Policy{
|
||||
IdleLifetime: 12 * time.Hour,
|
||||
AbsoluteLifetime: 24 * time.Hour,
|
||||
RememberIdleLifetime: 30 * 24 * time.Hour,
|
||||
RememberAbsoluteLifetime: 90 * 24 * time.Hour,
|
||||
}
|
||||
|
||||
// lifetimes picks the pair that applies to this session.
|
||||
func (p Policy) lifetimes(remember bool) (idle, absolute time.Duration) {
|
||||
if remember {
|
||||
return p.RememberIdleLifetime, p.RememberAbsoluteLifetime
|
||||
}
|
||||
return p.IdleLifetime, p.AbsoluteLifetime
|
||||
}
|
||||
|
||||
func (p Policy) validate() error {
|
||||
pairs := []struct {
|
||||
name string
|
||||
idle, absolute time.Duration
|
||||
}{
|
||||
{"normal", p.IdleLifetime, p.AbsoluteLifetime},
|
||||
{"remember-me", p.RememberIdleLifetime, p.RememberAbsoluteLifetime},
|
||||
}
|
||||
for _, pair := range pairs {
|
||||
if pair.idle <= 0 || pair.absolute <= 0 {
|
||||
return fmt.Errorf("auth: %s session lifetimes must be positive", pair.name)
|
||||
}
|
||||
// An absolute bound below the idle bound would make the idle window
|
||||
// unreachable, which is a configuration mistake rather than a policy.
|
||||
if pair.absolute < pair.idle {
|
||||
return fmt.Errorf("auth: %s absolute lifetime is shorter than its idle lifetime", pair.name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Store is the persistence the session lifecycle needs.
|
||||
//
|
||||
// An interface rather than a concrete type so that the lifecycle rules below
|
||||
// are testable without a database, and so this package does not depend on pgx.
|
||||
// The PostgreSQL implementation is PGStore, in store.go.
|
||||
//
|
||||
// Every method takes a token *hash*. No implementation ever receives a raw
|
||||
// token, which is what makes it structurally impossible to store one.
|
||||
type Store interface {
|
||||
// Create inserts the session and fills in the id the database assigned,
|
||||
// which is why it takes a pointer: the id is generated by the default on
|
||||
// the column, so the caller cannot know it beforehand.
|
||||
Create(ctx context.Context, s *Session) error
|
||||
FindByTokenHash(ctx context.Context, tokenHash string) (Session, error)
|
||||
Touch(ctx context.Context, id string, expiresAt, lastSeenAt time.Time) error
|
||||
Delete(ctx context.Context, id string) error
|
||||
DeleteByTokenHash(ctx context.Context, tokenHash string) error
|
||||
DeleteExpired(ctx context.Context, now time.Time) (int64, error)
|
||||
}
|
||||
|
||||
// Manager applies the session rules over a Store.
|
||||
//
|
||||
// It is the only place that turns a raw token into a hash, and the only place
|
||||
// that decides whether a session is still alive.
|
||||
type Manager struct {
|
||||
store Store
|
||||
policy Policy
|
||||
|
||||
// now is injectable so the expiry rules can be tested at a chosen instant
|
||||
// rather than by sleeping. Production always leaves it as time.Now.
|
||||
now func() time.Time
|
||||
|
||||
// slideThreshold avoids one UPDATE per request. The sliding deadline is
|
||||
// only pushed forward once the session has used up this fraction of its
|
||||
// idle window, so a burst of requests writes at most one row.
|
||||
slideThreshold float64
|
||||
}
|
||||
|
||||
// NewManager builds a Manager. An invalid policy is a programming error and is
|
||||
// reported here rather than at the first login.
|
||||
func NewManager(store Store, policy Policy) (*Manager, error) {
|
||||
if store == nil {
|
||||
return nil, errors.New("auth: session store is required")
|
||||
}
|
||||
if err := policy.validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Manager{store: store, policy: policy, now: time.Now, slideThreshold: 0.5}, nil
|
||||
}
|
||||
|
||||
// WithClock replaces the clock. For tests.
|
||||
func (m *Manager) WithClock(now func() time.Time) *Manager {
|
||||
if now != nil {
|
||||
m.now = now
|
||||
}
|
||||
return m
|
||||
}
|
||||
|
||||
// Policy is the lifetime policy in force.
|
||||
func (m *Manager) Policy() Policy { return m.policy }
|
||||
|
||||
// Issue creates a session for a user and returns the raw token exactly once.
|
||||
//
|
||||
// The token is the return value and is never stored: what reaches the database
|
||||
// is HashToken(token). The caller's only job is to put the raw token straight
|
||||
// into an HttpOnly cookie and then forget it — not log it, not echo it in a
|
||||
// JSON body, not put it in a URL.
|
||||
func (m *Manager) Issue(ctx context.Context, userID string, remember bool) (string, Session, error) {
|
||||
if userID == "" {
|
||||
return "", Session{}, errors.New("auth: user id is required")
|
||||
}
|
||||
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
return "", Session{}, err
|
||||
}
|
||||
|
||||
now := m.now().UTC()
|
||||
idle, absolute := m.policy.lifetimes(remember)
|
||||
s := Session{
|
||||
UserID: userID,
|
||||
TokenHash: HashToken(token),
|
||||
ExpiresAt: now.Add(idle),
|
||||
AbsoluteExpiresAt: now.Add(absolute),
|
||||
CreatedDate: now,
|
||||
LastSeenAt: now,
|
||||
}
|
||||
// The idle window can be the longer of the two only through a bad policy,
|
||||
// which validate() rejects; clamping anyway keeps the database CHECK
|
||||
// (expires_at <= absolute_expires_at) from being the thing that notices.
|
||||
if s.ExpiresAt.After(s.AbsoluteExpiresAt) {
|
||||
s.ExpiresAt = s.AbsoluteExpiresAt
|
||||
}
|
||||
|
||||
if err := m.store.Create(ctx, &s); err != nil {
|
||||
return "", Session{}, err
|
||||
}
|
||||
return token, s, nil
|
||||
}
|
||||
|
||||
// Authenticate resolves a raw token to a live session, sliding its expiry.
|
||||
//
|
||||
// It returns ErrSessionNotFound for an unknown token and ErrSessionExpired for
|
||||
// a dead one. Callers must answer the client identically in both cases: which
|
||||
// of the two it was tells an attacker whether a guessed token ever existed.
|
||||
//
|
||||
// An expired row is deleted as it is found, so a session that times out is
|
||||
// gone rather than waiting for the sweep.
|
||||
func (m *Manager) Authenticate(ctx context.Context, token string) (Session, error) {
|
||||
if token == "" {
|
||||
return Session{}, ErrEmptyToken
|
||||
}
|
||||
|
||||
hash := HashToken(token)
|
||||
s, err := m.store.FindByTokenHash(ctx, hash)
|
||||
if err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
|
||||
now := m.now().UTC()
|
||||
if s.IsExpired(now) {
|
||||
// Best effort: failing to delete does not make the session valid.
|
||||
_ = m.store.Delete(ctx, s.ID)
|
||||
return Session{}, ErrSessionExpired
|
||||
}
|
||||
|
||||
if err := m.slide(ctx, &s, now); err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// slide moves the sliding deadline forward, bounded by the absolute one.
|
||||
//
|
||||
// Only once the session is past slideThreshold of its idle window, so a page
|
||||
// that fires ten requests does not fire ten UPDATEs. The idle window is
|
||||
// recovered from the row rather than taken from the policy, so a session keeps
|
||||
// the lifetime it was issued under even if the policy changes underneath it.
|
||||
func (m *Manager) slide(ctx context.Context, s *Session, now time.Time) error {
|
||||
idle := s.ExpiresAt.Sub(s.LastSeenAt)
|
||||
if idle <= 0 {
|
||||
return nil
|
||||
}
|
||||
if now.Sub(s.LastSeenAt) < time.Duration(float64(idle)*m.slideThreshold) {
|
||||
return nil
|
||||
}
|
||||
|
||||
next := now.Add(idle)
|
||||
if next.After(s.AbsoluteExpiresAt) {
|
||||
next = s.AbsoluteExpiresAt
|
||||
}
|
||||
if err := m.store.Touch(ctx, s.ID, next, now); err != nil {
|
||||
return err
|
||||
}
|
||||
s.ExpiresAt = next
|
||||
s.LastSeenAt = now
|
||||
return nil
|
||||
}
|
||||
|
||||
// Lookup resolves a raw token without sliding the expiry or deleting anything.
|
||||
//
|
||||
// A read-only Authenticate, for callers that need to inspect a session without
|
||||
// treating the call as activity.
|
||||
func (m *Manager) Lookup(ctx context.Context, token string) (Session, error) {
|
||||
if token == "" {
|
||||
return Session{}, ErrEmptyToken
|
||||
}
|
||||
s, err := m.store.FindByTokenHash(ctx, HashToken(token))
|
||||
if err != nil {
|
||||
return Session{}, err
|
||||
}
|
||||
if s.IsExpired(m.now().UTC()) {
|
||||
return Session{}, ErrSessionExpired
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Revoke deletes the session behind a raw token. This is what logout calls.
|
||||
//
|
||||
// Deleting an already-absent session is not an error: logging out twice, or
|
||||
// logging out with a stale cookie, should succeed rather than fail loudly.
|
||||
func (m *Manager) Revoke(ctx context.Context, token string) error {
|
||||
if token == "" {
|
||||
return ErrEmptyToken
|
||||
}
|
||||
err := m.store.DeleteByTokenHash(ctx, HashToken(token))
|
||||
if errors.Is(err, ErrSessionNotFound) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// RevokeID deletes a session by its row id, for callers that already hold one.
|
||||
func (m *Manager) RevokeID(ctx context.Context, id string) error {
|
||||
if id == "" {
|
||||
return errors.New("auth: session id is required")
|
||||
}
|
||||
err := m.store.Delete(ctx, id)
|
||||
if errors.Is(err, ErrSessionNotFound) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
// Sweep deletes every session that is past either deadline, and reports how
|
||||
// many rows went. Authenticate already removes the expired sessions it meets;
|
||||
// this collects the ones nobody comes back for.
|
||||
func (m *Manager) Sweep(ctx context.Context) (int64, error) {
|
||||
return m.store.DeleteExpired(ctx, m.now().UTC())
|
||||
}
|
||||
488
go-api/internal/auth/session_test.go
Normal file
488
go-api/internal/auth/session_test.go
Normal file
@@ -0,0 +1,488 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// The session tests run against a real PostgreSQL database, because what they
|
||||
// are checking is largely the schema: the unique index that makes lookup work,
|
||||
// the CHECK that refuses a raw token, and the ON DELETE CASCADE that stops a
|
||||
// deleted user leaving a live session behind. None of that can be exercised
|
||||
// against an in-memory fake.
|
||||
//
|
||||
// Each test gets its own throwaway database, migrated but not seeded — the
|
||||
// seed fixture has nothing to say about sessions, and skipping it keeps these
|
||||
// tests fast. testutil skips rather than fails when PostgreSQL is absent.
|
||||
|
||||
type fixture struct {
|
||||
pool *pgxpool.Pool
|
||||
orgID string
|
||||
userID string
|
||||
store *PGStore
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func newFixture(t *testing.T, label string) *fixture {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
pool := testutil.Sandbox(t, label)
|
||||
testutil.ApplyAllMigrations(ctx, t, pool)
|
||||
|
||||
f := &fixture{pool: pool, ctx: ctx, store: NewPGStore(pool)}
|
||||
if err := pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
"Auth Test Org", "auth-test-org").Scan(&f.orgID); err != nil {
|
||||
t.Fatalf("create organization: %v", err)
|
||||
}
|
||||
f.userID = f.newUser(t, "session-owner@example.test")
|
||||
return f
|
||||
}
|
||||
|
||||
func (f *fixture) newUser(t *testing.T, email string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`INSERT INTO users (org_id, email, full_name, role)
|
||||
VALUES ($1::uuid, $2::citext, $3, 'admin') RETURNING id::text`,
|
||||
f.orgID, email, "Session Owner").Scan(&id); err != nil {
|
||||
t.Fatalf("create user %s: %v", email, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
func (f *fixture) sessionCount(t *testing.T) int {
|
||||
t.Helper()
|
||||
var n int
|
||||
if err := f.pool.QueryRow(f.ctx, `SELECT count(*)::int FROM sessions`).Scan(&n); err != nil {
|
||||
t.Fatalf("count sessions: %v", err)
|
||||
}
|
||||
return n
|
||||
}
|
||||
|
||||
// manager builds a Manager over the fixture's store with a clock the test drives.
|
||||
func (f *fixture) manager(t *testing.T, p Policy, clock *time.Time) *Manager {
|
||||
t.Helper()
|
||||
m, err := NewManager(f.store, p)
|
||||
if err != nil {
|
||||
t.Fatalf("NewManager: %v", err)
|
||||
}
|
||||
return m.WithClock(func() time.Time { return *clock })
|
||||
}
|
||||
|
||||
// shortPolicy keeps the arithmetic in these tests small and legible. The
|
||||
// production values are asserted separately, in TestDefaultPolicy.
|
||||
var shortPolicy = Policy{
|
||||
IdleLifetime: time.Hour,
|
||||
AbsoluteLifetime: 3 * time.Hour,
|
||||
RememberIdleLifetime: 24 * time.Hour,
|
||||
RememberAbsoluteLifetime: 72 * time.Hour,
|
||||
}
|
||||
|
||||
// 9. A session is created, and what lands in the database is the hash.
|
||||
func TestSessionCreation(t *testing.T) {
|
||||
f := newFixture(t, "session_create")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
|
||||
if token == "" {
|
||||
t.Fatal("Issue returned an empty token")
|
||||
}
|
||||
if sess.ID == "" {
|
||||
t.Error("the session was not given an id")
|
||||
}
|
||||
if sess.UserID != f.userID {
|
||||
t.Errorf("UserID = %q, want %q", sess.UserID, f.userID)
|
||||
}
|
||||
if sess.TokenHash != HashToken(token) {
|
||||
t.Error("the session's token hash is not the hash of the returned token")
|
||||
}
|
||||
if sess.TokenHash == token {
|
||||
t.Fatal("the raw token was stored as the hash")
|
||||
}
|
||||
if want := now.Add(shortPolicy.IdleLifetime); !sess.ExpiresAt.Equal(want) {
|
||||
t.Errorf("ExpiresAt = %v, want %v", sess.ExpiresAt, want)
|
||||
}
|
||||
if want := now.Add(shortPolicy.AbsoluteLifetime); !sess.AbsoluteExpiresAt.Equal(want) {
|
||||
t.Errorf("AbsoluteExpiresAt = %v, want %v", sess.AbsoluteExpiresAt, want)
|
||||
}
|
||||
|
||||
// The row itself: exactly one, holding the hash and never the token.
|
||||
var stored string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT token_hash FROM sessions WHERE id = $1::uuid`, sess.ID).Scan(&stored); err != nil {
|
||||
t.Fatalf("read the stored session: %v", err)
|
||||
}
|
||||
if stored != HashToken(token) {
|
||||
t.Error("the stored token_hash is not the hash of the token")
|
||||
}
|
||||
|
||||
// Nothing anywhere in the table equals the token. This is the property the
|
||||
// whole design exists for, so it is asserted against the database and not
|
||||
// against the struct.
|
||||
var leaked int
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT count(*)::int FROM sessions WHERE token_hash = $1::text`, token).Scan(&leaked); err != nil {
|
||||
t.Fatalf("scan for a leaked token: %v", err)
|
||||
}
|
||||
if leaked != 0 {
|
||||
t.Fatal("the raw token is present in the database")
|
||||
}
|
||||
|
||||
// Remember Me gets the longer pair of deadlines.
|
||||
_, remembered, err := m.Issue(f.ctx, f.userID, true)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue with remember: %v", err)
|
||||
}
|
||||
if want := now.Add(shortPolicy.RememberIdleLifetime); !remembered.ExpiresAt.Equal(want) {
|
||||
t.Errorf("Remember Me ExpiresAt = %v, want %v", remembered.ExpiresAt, want)
|
||||
}
|
||||
if want := now.Add(shortPolicy.RememberAbsoluteLifetime); !remembered.AbsoluteExpiresAt.Equal(want) {
|
||||
t.Errorf("Remember Me AbsoluteExpiresAt = %v, want %v", remembered.AbsoluteExpiresAt, want)
|
||||
}
|
||||
}
|
||||
|
||||
// 10. A live session is found by the token, and only by the right token.
|
||||
func TestSessionLookup(t *testing.T) {
|
||||
f := newFixture(t, "session_lookup")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, issued, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
|
||||
got, err := m.Authenticate(f.ctx, token)
|
||||
if err != nil {
|
||||
t.Fatalf("Authenticate: %v", err)
|
||||
}
|
||||
if got.ID != issued.ID || got.UserID != f.userID {
|
||||
t.Errorf("Authenticate returned session %q for user %q, want %q / %q",
|
||||
got.ID, got.UserID, issued.ID, f.userID)
|
||||
}
|
||||
|
||||
// Lookup is the read-only form and must agree.
|
||||
if looked, err := m.Lookup(f.ctx, token); err != nil || looked.ID != issued.ID {
|
||||
t.Errorf("Lookup = %q, %v; want %q, nil", looked.ID, err, issued.ID)
|
||||
}
|
||||
|
||||
// A different, valid-looking token must not resolve. This is the case the
|
||||
// unique index and the hash lookup exist to make hopeless.
|
||||
other, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, other); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("an unknown token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
|
||||
// A malformed token must be refused without a round trip, and an empty one
|
||||
// must be refused before it is hashed at all.
|
||||
if _, err := m.Authenticate(f.ctx, "not-a-real-token"); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("a malformed token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, ""); !errors.Is(err, ErrEmptyToken) {
|
||||
t.Errorf("an empty token returned %v, want ErrEmptyToken", err)
|
||||
}
|
||||
|
||||
// The store is keyed by hash and refuses a raw token outright, so a caller
|
||||
// that forgets to hash gets an error rather than a silent miss.
|
||||
if _, err := f.store.FindByTokenHash(f.ctx, token); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("FindByTokenHash with a raw token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 11. An expired session is rejected, and is cleaned up as it is found.
|
||||
func TestExpiredSessionIsRejected(t *testing.T) {
|
||||
f := newFixture(t, "session_expiry")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
|
||||
// One second before the deadline it is still good.
|
||||
now = sess.ExpiresAt.Add(-time.Second)
|
||||
if _, err := m.Authenticate(f.ctx, token); err != nil {
|
||||
t.Fatalf("a session one second from expiry was rejected: %v", err)
|
||||
}
|
||||
|
||||
// Exactly at the deadline it is not. The boundary is closed, not open: a
|
||||
// session that expires at 12:00 is dead at 12:00.
|
||||
fresh, err := m.Lookup(f.ctx, token)
|
||||
if err != nil {
|
||||
t.Fatalf("Lookup: %v", err)
|
||||
}
|
||||
now = fresh.ExpiresAt
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionExpired) {
|
||||
t.Fatalf("an expired session returned %v, want ErrSessionExpired", err)
|
||||
}
|
||||
|
||||
// Finding it expired removes it, so the row does not linger until a sweep.
|
||||
if n := f.sessionCount(t); n != 0 {
|
||||
t.Errorf("%d sessions remain after an expired one was authenticated, want 0", n)
|
||||
}
|
||||
// And the second attempt cannot tell the client anything different.
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("re-authenticating a swept session returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// The absolute ceiling is what stops a sliding window being slid forever.
|
||||
func TestSessionCannotOutliveItsAbsoluteDeadline(t *testing.T) {
|
||||
f := newFixture(t, "session_absolute")
|
||||
start := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
now := start
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
ceiling := sess.AbsoluteExpiresAt
|
||||
|
||||
// Keep using the session steadily, well inside the idle window each time,
|
||||
// right up to the ceiling. The sliding deadline must never cross it.
|
||||
for _, at := range []time.Duration{50 * time.Minute, 105 * time.Minute, 160 * time.Minute, 175 * time.Minute} {
|
||||
now = start.Add(at)
|
||||
got, err := m.Authenticate(f.ctx, token)
|
||||
if err != nil {
|
||||
t.Fatalf("Authenticate at +%v: %v", at, err)
|
||||
}
|
||||
if got.ExpiresAt.After(ceiling) {
|
||||
t.Fatalf("at +%v the sliding deadline %v passed the absolute ceiling %v",
|
||||
at, got.ExpiresAt, ceiling)
|
||||
}
|
||||
if !got.AbsoluteExpiresAt.Equal(ceiling) {
|
||||
t.Fatalf("at +%v the absolute ceiling moved to %v, want %v", at, got.AbsoluteExpiresAt, ceiling)
|
||||
}
|
||||
}
|
||||
|
||||
// At the ceiling the session is over, however recently it was used.
|
||||
now = ceiling
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionExpired) {
|
||||
t.Fatalf("at the absolute ceiling Authenticate returned %v, want ErrSessionExpired", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 12. Revoking a session deletes it — and revoking twice is not an error,
|
||||
// because logging out with a stale cookie should succeed.
|
||||
func TestSessionDeletion(t *testing.T) {
|
||||
f := newFixture(t, "session_delete")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
token, sess, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
if err := m.Revoke(f.ctx, token); err != nil {
|
||||
t.Fatalf("Revoke: %v", err)
|
||||
}
|
||||
if n := f.sessionCount(t); n != 0 {
|
||||
t.Errorf("%d sessions remain after revocation, want 0", n)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, token); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("a revoked token returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
if err := m.Revoke(f.ctx, token); err != nil {
|
||||
t.Errorf("revoking twice returned %v, want nil", err)
|
||||
}
|
||||
if err := m.RevokeID(f.ctx, sess.ID); err != nil {
|
||||
t.Errorf("revoking an absent session by id returned %v, want nil", err)
|
||||
}
|
||||
|
||||
// The store's own contract is stricter: it reports what it did.
|
||||
if _, _, err := m.Issue(f.ctx, f.userID, false); err != nil {
|
||||
t.Fatalf("Issue: %v", err)
|
||||
}
|
||||
var id string
|
||||
if err := f.pool.QueryRow(f.ctx, `SELECT id::text FROM sessions LIMIT 1`).Scan(&id); err != nil {
|
||||
t.Fatalf("read the session id: %v", err)
|
||||
}
|
||||
if err := f.store.Delete(f.ctx, id); err != nil {
|
||||
t.Fatalf("store.Delete: %v", err)
|
||||
}
|
||||
if err := f.store.Delete(f.ctx, id); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("deleting an absent session returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 13. Deleting a user deletes their sessions. Without this, a removed account
|
||||
// keeps working until its cookie happens to expire.
|
||||
func TestUserDeletionCascadesToSessions(t *testing.T) {
|
||||
f := newFixture(t, "session_cascade")
|
||||
now := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
// A second user, so the test can prove the cascade removes one user's
|
||||
// sessions and leaves the other's alone.
|
||||
otherID := f.newUser(t, "other-user@example.test")
|
||||
|
||||
doomedToken, _, err := m.Issue(f.ctx, f.userID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue for the doomed user: %v", err)
|
||||
}
|
||||
if _, _, err := m.Issue(f.ctx, f.userID, true); err != nil {
|
||||
t.Fatalf("second Issue for the doomed user: %v", err)
|
||||
}
|
||||
survivingToken, _, err := m.Issue(f.ctx, otherID, false)
|
||||
if err != nil {
|
||||
t.Fatalf("Issue for the surviving user: %v", err)
|
||||
}
|
||||
if n := f.sessionCount(t); n != 3 {
|
||||
t.Fatalf("%d sessions before the delete, want 3", n)
|
||||
}
|
||||
|
||||
if _, err := f.pool.Exec(f.ctx, `DELETE FROM users WHERE id = $1::uuid`, f.userID); err != nil {
|
||||
// A RESTRICT foreign key would fail here, which is exactly the design
|
||||
// this test rules out.
|
||||
t.Fatalf("delete the user: %v", err)
|
||||
}
|
||||
|
||||
if n := f.sessionCount(t); n != 1 {
|
||||
t.Fatalf("%d sessions after deleting one of two users, want 1", n)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, doomedToken); !errors.Is(err, ErrSessionNotFound) {
|
||||
t.Errorf("a deleted user's session returned %v, want ErrSessionNotFound", err)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, survivingToken); err != nil {
|
||||
t.Errorf("the other user's session was destroyed too: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Sweep collects the sessions nobody comes back for, by either deadline.
|
||||
func TestSweepDeletesExpiredSessions(t *testing.T) {
|
||||
f := newFixture(t, "session_sweep")
|
||||
start := time.Date(2026, 8, 22, 12, 0, 0, 0, time.UTC)
|
||||
now := start
|
||||
m := f.manager(t, shortPolicy, &now)
|
||||
|
||||
if _, _, err := m.Issue(f.ctx, f.userID, false); err != nil { // dies at +1h
|
||||
t.Fatalf("Issue short: %v", err)
|
||||
}
|
||||
longToken, _, err := m.Issue(f.ctx, f.userID, true) // dies at +24h
|
||||
if err != nil {
|
||||
t.Fatalf("Issue long: %v", err)
|
||||
}
|
||||
|
||||
now = start.Add(90 * time.Minute)
|
||||
n, err := m.Sweep(f.ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Sweep: %v", err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Errorf("Sweep removed %d sessions, want 1", n)
|
||||
}
|
||||
if _, err := m.Authenticate(f.ctx, longToken); err != nil {
|
||||
t.Errorf("Sweep removed a live session: %v", err)
|
||||
}
|
||||
|
||||
now = start.Add(25 * time.Hour)
|
||||
if n, err = m.Sweep(f.ctx); err != nil || n != 1 {
|
||||
t.Errorf("second Sweep removed %d sessions (err %v), want 1", n, err)
|
||||
}
|
||||
if got := f.sessionCount(t); got != 0 {
|
||||
t.Errorf("%d sessions remain after the sweep, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
// The database is the last line of defence against storing a raw token: the
|
||||
// CHECK constraint refuses anything that is not a SHA-256 hex digest, even if
|
||||
// the Go guard were bypassed.
|
||||
func TestDatabaseRefusesARawToken(t *testing.T) {
|
||||
f := newFixture(t, "session_rawtoken")
|
||||
now := time.Now().UTC()
|
||||
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
// Through the store: a Go error naming the mistake.
|
||||
err = f.store.Create(f.ctx, &Session{
|
||||
UserID: f.userID, TokenHash: token,
|
||||
ExpiresAt: now.Add(time.Hour), AbsoluteExpiresAt: now.Add(2 * time.Hour),
|
||||
CreatedDate: now, LastSeenAt: now,
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("the store accepted a raw token as a token hash")
|
||||
}
|
||||
|
||||
// Straight past the store, in SQL: the constraint still refuses it.
|
||||
_, err = f.pool.Exec(f.ctx,
|
||||
`INSERT INTO sessions (user_id, token_hash, expires_at, absolute_expires_at)
|
||||
VALUES ($1::uuid, $2::text, now() + interval '1 hour', now() + interval '2 hours')`,
|
||||
f.userID, token)
|
||||
if err == nil {
|
||||
t.Fatal("the sessions table accepted a raw token; the CHECK constraint is not doing its job")
|
||||
}
|
||||
|
||||
// The same insert with a proper hash succeeds, so the constraint is not
|
||||
// simply rejecting everything.
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO sessions (user_id, token_hash, expires_at, absolute_expires_at)
|
||||
VALUES ($1::uuid, $2::text, now() + interval '1 hour', now() + interval '2 hours')`,
|
||||
f.userID, HashToken(token)); err != nil {
|
||||
t.Fatalf("a well-formed session was refused: %v", err)
|
||||
}
|
||||
|
||||
// And the same hash cannot be stored twice: UNIQUE is what makes a token
|
||||
// identify exactly one session.
|
||||
if _, err := f.pool.Exec(f.ctx,
|
||||
`INSERT INTO sessions (user_id, token_hash, expires_at, absolute_expires_at)
|
||||
VALUES ($1::uuid, $2::text, now() + interval '1 hour', now() + interval '2 hours')`,
|
||||
f.userID, HashToken(token)); err == nil {
|
||||
t.Fatal("two sessions were stored with the same token hash")
|
||||
}
|
||||
}
|
||||
|
||||
// The lifetimes are a decision, so they are asserted rather than assumed.
|
||||
func TestDefaultPolicy(t *testing.T) {
|
||||
p := DefaultPolicy
|
||||
if p.IdleLifetime != 12*time.Hour {
|
||||
t.Errorf("IdleLifetime = %v, want 12h", p.IdleLifetime)
|
||||
}
|
||||
if p.RememberIdleLifetime != 30*24*time.Hour {
|
||||
t.Errorf("RememberIdleLifetime = %v, want 720h (30 days)", p.RememberIdleLifetime)
|
||||
}
|
||||
// The point of the absolute bound: it must exist and must exceed the
|
||||
// window it caps, or a session could live forever.
|
||||
if p.AbsoluteLifetime <= 0 || p.RememberAbsoluteLifetime <= 0 {
|
||||
t.Fatal("an absolute lifetime is unset; a session could live forever")
|
||||
}
|
||||
if err := p.validate(); err != nil {
|
||||
t.Errorf("the default policy is not valid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewManagerRejectsBadInput(t *testing.T) {
|
||||
if _, err := NewManager(nil, DefaultPolicy); err == nil {
|
||||
t.Error("NewManager accepted a nil store")
|
||||
}
|
||||
bad := map[string]Policy{
|
||||
"zero": {},
|
||||
"absolute below idle": {IdleLifetime: 2 * time.Hour, AbsoluteLifetime: time.Hour, RememberIdleLifetime: time.Hour, RememberAbsoluteLifetime: time.Hour},
|
||||
"remember-me unset": {IdleLifetime: time.Hour, AbsoluteLifetime: time.Hour},
|
||||
"negative idle lifetime": {IdleLifetime: -time.Hour, AbsoluteLifetime: time.Hour, RememberIdleLifetime: time.Hour, RememberAbsoluteLifetime: time.Hour},
|
||||
}
|
||||
for name, p := range bad {
|
||||
if _, err := NewManager(NewPGStore(nil), p); err == nil {
|
||||
t.Errorf("NewManager accepted the %s policy", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
204
go-api/internal/auth/store.go
Normal file
204
go-api/internal/auth/store.go
Normal file
@@ -0,0 +1,204 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
)
|
||||
|
||||
// Querier is satisfied by *pgxpool.Pool and by pgx.Tx, so every method below
|
||||
// works inside or outside a transaction.
|
||||
//
|
||||
// Declared here rather than imported from internal/repo: that package's
|
||||
// Querier is identical, but this one keeps the authentication foundation
|
||||
// independent of the resource/descriptor layer, which it otherwise shares
|
||||
// nothing with. Go interfaces are structural, so both are satisfied by the
|
||||
// same values.
|
||||
type Querier interface {
|
||||
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
|
||||
QueryRow(ctx context.Context, sql string, args ...any) pgx.Row
|
||||
Exec(ctx context.Context, sql string, args ...any) (pgconn.CommandTag, error)
|
||||
}
|
||||
|
||||
// PGStore is the sessions table.
|
||||
//
|
||||
// Every statement here is a constant string with bind parameters. Nothing —
|
||||
// not an id, not a token hash, not a timestamp — is ever formatted into SQL.
|
||||
// There is no identifier taken from a caller, so there is nothing to quote and
|
||||
// nothing to escape.
|
||||
type PGStore struct {
|
||||
db Querier
|
||||
}
|
||||
|
||||
// NewPGStore builds the store over a pool or a transaction.
|
||||
func NewPGStore(db Querier) *PGStore { return &PGStore{db: db} }
|
||||
|
||||
// Compile-time check that the persistence layer satisfies the lifecycle's
|
||||
// expectations. If Store gains a method, this line is where it is noticed.
|
||||
var _ Store = (*PGStore)(nil)
|
||||
|
||||
// sessionColumns is the projection every read below shares, in the order the
|
||||
// scan expects.
|
||||
const sessionColumns = `id::text, user_id::text, token_hash,
|
||||
expires_at, absolute_expires_at, created_date, last_seen_at`
|
||||
|
||||
// Create inserts a session.
|
||||
//
|
||||
// The id is left to the database's gen_random_uuid() default when the caller
|
||||
// did not choose one, and returned so the caller's Session is complete.
|
||||
// created_date and last_seen_at come from the caller rather than now(), so the
|
||||
// row agrees with the deadlines the Manager computed from the same instant.
|
||||
func (s *PGStore) Create(ctx context.Context, sess *Session) error {
|
||||
if sess.UserID == "" {
|
||||
return errors.New("auth: session user id is required")
|
||||
}
|
||||
// The last line of defence against writing a raw token to disk. The
|
||||
// database CHECK enforces the same shape; this turns it into a Go error
|
||||
// naming the actual mistake instead of a constraint violation.
|
||||
if !IsTokenHash(sess.TokenHash) {
|
||||
return errors.New("auth: session token_hash is not a SHA-256 hex digest")
|
||||
}
|
||||
|
||||
const q = `INSERT INTO sessions
|
||||
(id, user_id, token_hash, expires_at, absolute_expires_at, created_date, last_seen_at)
|
||||
VALUES
|
||||
(COALESCE($1::uuid, gen_random_uuid()), $2::uuid, $3::text,
|
||||
$4::timestamptz, $5::timestamptz, $6::timestamptz, $7::timestamptz)
|
||||
RETURNING id::text`
|
||||
|
||||
var id *string
|
||||
if sess.ID != "" {
|
||||
id = &sess.ID
|
||||
}
|
||||
if err := s.db.QueryRow(ctx, q, id, sess.UserID, sess.TokenHash,
|
||||
sess.ExpiresAt, sess.AbsoluteExpiresAt, sess.CreatedDate, sess.LastSeenAt,
|
||||
).Scan(&sess.ID); err != nil {
|
||||
return fmt.Errorf("auth: create session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindByTokenHash reads one session by the hash of its token.
|
||||
//
|
||||
// The parameter is a hash, never a token: the Manager hashes before it calls
|
||||
// here, so a raw secret never reaches the query layer at all. Returns
|
||||
// ErrSessionNotFound when no row matches, which callers must not distinguish
|
||||
// from an expired session when answering a client.
|
||||
//
|
||||
// Expiry is deliberately not filtered in SQL. The caller decides what an
|
||||
// expired row means — Authenticate deletes it, a diagnostic might report it —
|
||||
// and a WHERE clause here would collapse "revoked" and "timed out" into one
|
||||
// indistinguishable answer at the wrong layer.
|
||||
func (s *PGStore) FindByTokenHash(ctx context.Context, tokenHash string) (Session, error) {
|
||||
if tokenHash == "" {
|
||||
return Session{}, ErrEmptyToken
|
||||
}
|
||||
if !IsTokenHash(tokenHash) {
|
||||
// A value of the wrong shape cannot match any row, and querying with
|
||||
// it would be an unnecessary round trip on every malformed cookie.
|
||||
return Session{}, ErrSessionNotFound
|
||||
}
|
||||
|
||||
const q = `SELECT ` + sessionColumns + ` FROM sessions WHERE token_hash = $1::text`
|
||||
|
||||
var out Session
|
||||
err := s.db.QueryRow(ctx, q, tokenHash).Scan(
|
||||
&out.ID, &out.UserID, &out.TokenHash,
|
||||
&out.ExpiresAt, &out.AbsoluteExpiresAt, &out.CreatedDate, &out.LastSeenAt)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return Session{}, ErrSessionNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return Session{}, fmt.Errorf("auth: find session: %w", err)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// Touch moves the sliding deadline and records the activity.
|
||||
//
|
||||
// The UPDATE is guarded by `expires_at <= absolute_expires_at` in the database
|
||||
// CHECK; the Manager clamps before calling, so a violation here would mean a
|
||||
// bug rather than a race.
|
||||
func (s *PGStore) Touch(ctx context.Context, id string, expiresAt, lastSeenAt time.Time) error {
|
||||
if id == "" {
|
||||
return errors.New("auth: session id is required")
|
||||
}
|
||||
|
||||
const q = `UPDATE sessions
|
||||
SET expires_at = $2::timestamptz, last_seen_at = $3::timestamptz
|
||||
WHERE id = $1::uuid`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, id, expiresAt, lastSeenAt)
|
||||
if err != nil {
|
||||
return fmt.Errorf("auth: touch session: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
// The session was revoked between the read and this write. Reporting
|
||||
// it as absent is honest; the caller treats it as a failed lookup.
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Delete removes one session by id. Returns ErrSessionNotFound if there was
|
||||
// nothing to remove; Manager.RevokeID absorbs that, because logging out of a
|
||||
// session that is already gone is a success.
|
||||
func (s *PGStore) Delete(ctx context.Context, id string) error {
|
||||
if id == "" {
|
||||
return errors.New("auth: session id is required")
|
||||
}
|
||||
|
||||
const q = `DELETE FROM sessions WHERE id = $1::uuid`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("auth: delete session: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteByTokenHash removes one session by the hash of its token. This is the
|
||||
// logout path: the cookie is all the client has.
|
||||
func (s *PGStore) DeleteByTokenHash(ctx context.Context, tokenHash string) error {
|
||||
if tokenHash == "" {
|
||||
return ErrEmptyToken
|
||||
}
|
||||
if !IsTokenHash(tokenHash) {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
|
||||
const q = `DELETE FROM sessions WHERE token_hash = $1::text`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, tokenHash)
|
||||
if err != nil {
|
||||
return fmt.Errorf("auth: delete session by token: %w", err)
|
||||
}
|
||||
if tag.RowsAffected() == 0 {
|
||||
return ErrSessionNotFound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteExpired removes every session past either deadline and reports the
|
||||
// count.
|
||||
//
|
||||
// Both deadlines are tested. Filtering on expires_at alone would leave behind
|
||||
// a session whose sliding window is still open but whose absolute ceiling has
|
||||
// passed — precisely the row the absolute bound exists to kill.
|
||||
func (s *PGStore) DeleteExpired(ctx context.Context, now time.Time) (int64, error) {
|
||||
const q = `DELETE FROM sessions
|
||||
WHERE expires_at <= $1::timestamptz OR absolute_expires_at <= $1::timestamptz`
|
||||
|
||||
tag, err := s.db.Exec(ctx, q, now)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("auth: delete expired sessions: %w", err)
|
||||
}
|
||||
return tag.RowsAffected(), nil
|
||||
}
|
||||
74
go-api/internal/auth/token.go
Normal file
74
go-api/internal/auth/token.go
Normal file
@@ -0,0 +1,74 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"regexp"
|
||||
)
|
||||
|
||||
// ErrEmptyToken is returned when a token is required and none was supplied.
|
||||
// It is deliberately distinct from "no such session": an absent cookie is a
|
||||
// different situation from a cookie that no longer matches a row.
|
||||
var ErrEmptyToken = errors.New("auth: session token is empty")
|
||||
|
||||
// TokenBytes is the entropy behind a session token.
|
||||
//
|
||||
// 32 bytes — 256 bits — is the size of the SHA-256 output it is hashed to, so
|
||||
// nothing is wasted at either end, and it puts guessing a live session far
|
||||
// beyond reach: an attacker who could test a billion candidates a second would
|
||||
// still need on the order of 10^60 years.
|
||||
//
|
||||
// This is a raw byte count, not a character count. The encoded token is 43
|
||||
// characters of base64url.
|
||||
const TokenBytes = 32
|
||||
|
||||
// tokenHashPattern is the exact shape stored in sessions.token_hash, and the
|
||||
// same pattern the sessions_token_hash_sha256 CHECK constraint enforces in
|
||||
// migration 000004. Validating here turns a database constraint violation into
|
||||
// a clear Go error at the point the mistake was made.
|
||||
var tokenHashPattern = regexp.MustCompile(`^[0-9a-f]{64}$`)
|
||||
|
||||
// GenerateToken returns a new, cryptographically random session token.
|
||||
//
|
||||
// The returned string is the secret itself. It is what goes into the HttpOnly
|
||||
// cookie and it must never be written to the database, to a log, or to an
|
||||
// error message. Only its hash is persisted — see HashToken.
|
||||
//
|
||||
// base64url without padding, so the value is safe in a cookie, a header and a
|
||||
// URL without escaping, and contains no '=' to be mangled by a cookie parser.
|
||||
func GenerateToken() (string, error) {
|
||||
buf := make([]byte, TokenBytes)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
// There is no fallback. math/rand here would produce tokens an
|
||||
// attacker can predict from a handful of observed sessions.
|
||||
return "", fmt.Errorf("auth: read random bytes: %w", err)
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
// HashToken returns the lowercase hex SHA-256 of a session token.
|
||||
//
|
||||
// This is what the database stores. A plain hash — not argon2 — is the right
|
||||
// choice here and the wrong one for a password, and the difference is entropy:
|
||||
// a session token is 256 uniformly random bits, so there is no dictionary to
|
||||
// run against it and no work factor worth paying on every single request. A
|
||||
// password is chosen by a human and needs argon2id precisely because it is not.
|
||||
//
|
||||
// The function is pure and deterministic: the same token always hashes to the
|
||||
// same string, which is what makes lookup by hash possible at all.
|
||||
func HashToken(token string) string {
|
||||
sum := sha256.Sum256([]byte(token))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// IsTokenHash reports whether s has the shape HashToken produces.
|
||||
//
|
||||
// Used to catch the one mistake that would be catastrophic and silent: passing
|
||||
// a raw token where a hash is expected, and storing the secret in plaintext.
|
||||
// A raw token is base64url and contains characters outside [0-9a-f], or is the
|
||||
// wrong length, so it always fails this test.
|
||||
func IsTokenHash(s string) bool { return tokenHashPattern.MatchString(s) }
|
||||
158
go-api/internal/auth/token_test.go
Normal file
158
go-api/internal/auth/token_test.go
Normal file
@@ -0,0 +1,158 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// 6. Session tokens carry real entropy from crypto/rand.
|
||||
//
|
||||
// Randomness cannot be proved by a test, so this asserts the properties whose
|
||||
// absence would mean the generator is broken: the full 256 bits are present,
|
||||
// the output is not a constant, and the bytes are not all the same value —
|
||||
// which is what a zeroed or unseeded buffer looks like.
|
||||
func TestGenerateTokenIsCryptographicallyRandom(t *testing.T) {
|
||||
const runs = 512
|
||||
|
||||
seen := make(map[string]struct{}, runs)
|
||||
bitsSet := make([]int, 8*TokenBytes) // how often each bit position was 1
|
||||
|
||||
for i := 0; i < runs; i++ {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
raw, err := base64.RawURLEncoding.DecodeString(token)
|
||||
if err != nil {
|
||||
t.Fatalf("token is not base64url: %v", err)
|
||||
}
|
||||
if len(raw) != TokenBytes {
|
||||
t.Fatalf("token decodes to %d bytes, want %d", len(raw), TokenBytes)
|
||||
}
|
||||
// A cookie value must survive a round trip untouched: no '=' padding,
|
||||
// no '+' or '/' to be re-encoded.
|
||||
if strings.ContainsAny(token, "=+/") {
|
||||
t.Fatalf("token contains a character that is unsafe in a cookie or URL")
|
||||
}
|
||||
|
||||
if _, dup := seen[token]; dup {
|
||||
t.Fatalf("GenerateToken returned a duplicate within %d calls", runs)
|
||||
}
|
||||
seen[token] = struct{}{}
|
||||
|
||||
for bit := 0; bit < 8*TokenBytes; bit++ {
|
||||
if raw[bit/8]&(1<<(bit%8)) != 0 {
|
||||
bitsSet[bit]++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Each bit should be 1 about half the time. A bit that is *always* 0 or
|
||||
// always 1 across 512 draws has a chance of roughly 2^-511 of being random
|
||||
// and is far more likely a stuck generator. The bound is deliberately
|
||||
// loose — this is a smoke test for a broken source, not a statistical
|
||||
// suite, and it must never flake.
|
||||
for bit, count := range bitsSet {
|
||||
if count == 0 || count == runs {
|
||||
t.Errorf("bit %d was constant across %d tokens; the entropy source is broken", bit, runs)
|
||||
}
|
||||
}
|
||||
|
||||
// 256 bits is the size that makes guessing a live session hopeless.
|
||||
if TokenBytes < 32 {
|
||||
t.Errorf("TokenBytes = %d, want at least 32", TokenBytes)
|
||||
}
|
||||
}
|
||||
|
||||
// 7. Two generated tokens differ.
|
||||
func TestGenerateTokenReturnsDistinctValues(t *testing.T) {
|
||||
a, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("first: %v", err)
|
||||
}
|
||||
b, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("second: %v", err)
|
||||
}
|
||||
if a == b {
|
||||
t.Fatal("two consecutive tokens were identical")
|
||||
}
|
||||
if HashToken(a) == HashToken(b) {
|
||||
t.Fatal("two distinct tokens hashed to the same value")
|
||||
}
|
||||
}
|
||||
|
||||
// 8. Hashing is deterministic, and is genuinely SHA-256 rather than something
|
||||
// that merely looks like it.
|
||||
func TestHashTokenIsDeterministicSHA256(t *testing.T) {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
first, second := HashToken(token), HashToken(token)
|
||||
if first != second {
|
||||
t.Fatal("hashing the same token twice produced different values")
|
||||
}
|
||||
|
||||
// Checked against the standard library directly: lookup only works if the
|
||||
// value stored is exactly this.
|
||||
want := sha256.Sum256([]byte(token))
|
||||
if first != hex.EncodeToString(want[:]) {
|
||||
t.Fatal("HashToken does not agree with crypto/sha256")
|
||||
}
|
||||
|
||||
if len(first) != 64 {
|
||||
t.Fatalf("hash is %d characters, want 64 hex characters", len(first))
|
||||
}
|
||||
if first != strings.ToLower(first) {
|
||||
t.Error("hash is not lowercase; the database CHECK requires lowercase hex")
|
||||
}
|
||||
// The stored value must not be the secret.
|
||||
if strings.Contains(first, token) || first == token {
|
||||
t.Error("the hash contains the token")
|
||||
}
|
||||
if HashToken(token+"x") == first {
|
||||
t.Error("a different token hashed to the same value")
|
||||
}
|
||||
|
||||
// The empty string has a hash too — that is a property of SHA-256, not a
|
||||
// licence to store one. Guarding against an empty token is the Manager's
|
||||
// job, and TestManagerRejectsEmptyToken covers it.
|
||||
if HashToken("") == "" {
|
||||
t.Error("HashToken returned an empty string")
|
||||
}
|
||||
}
|
||||
|
||||
// IsTokenHash is the guard that stops a raw token being written where a hash
|
||||
// belongs, so it must reject every raw token and accept every real hash.
|
||||
func TestIsTokenHash(t *testing.T) {
|
||||
token, err := GenerateToken()
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateToken: %v", err)
|
||||
}
|
||||
|
||||
if !IsTokenHash(HashToken(token)) {
|
||||
t.Error("a real hash was not recognised as one")
|
||||
}
|
||||
if IsTokenHash(token) {
|
||||
t.Error("a raw token was accepted as a hash; this is the check that prevents storing the secret")
|
||||
}
|
||||
|
||||
for name, s := range map[string]string{
|
||||
"empty": "",
|
||||
"too short": strings.Repeat("a", 63),
|
||||
"too long": strings.Repeat("a", 65),
|
||||
"uppercase": strings.ToUpper(HashToken(token)),
|
||||
"non-hex": strings.Repeat("g", 64),
|
||||
"trailing": HashToken(token) + "\n",
|
||||
} {
|
||||
if IsTokenHash(s) {
|
||||
t.Errorf("%s was accepted as a token hash", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
142
go-api/internal/auth/users.go
Normal file
142
go-api/internal/auth/users.go
Normal file
@@ -0,0 +1,142 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
)
|
||||
|
||||
// ErrUserNotFound means no user matched. Callers authenticating someone must
|
||||
// answer this identically to a wrong password — see the note on Credentials.
|
||||
var ErrUserNotFound = errors.New("auth: user not found")
|
||||
|
||||
// StatusActive is the only users.status that may hold a session. The column's
|
||||
// CHECK constraint (migration 000001) permits 'active' and 'suspended'.
|
||||
const StatusActive = "active"
|
||||
|
||||
// User is the part of a users row that authentication needs.
|
||||
//
|
||||
// Note what is absent: preferences, timestamps, legacy ids. This is not a
|
||||
// general user model — the resource layer already has one — it is the set of
|
||||
// facts required to answer "may this person hold a session, and whose data do
|
||||
// they see".
|
||||
type User struct {
|
||||
ID string
|
||||
OrgID string
|
||||
Email string
|
||||
FullName string
|
||||
Role string
|
||||
AccountType string
|
||||
Status string
|
||||
|
||||
// PasswordHash is empty when the user has never had a password set.
|
||||
// migration 000001 leaves the column nullable and NULL, and the seeded
|
||||
// demo user is in exactly that state until `setpassword` is run.
|
||||
PasswordHash string
|
||||
}
|
||||
|
||||
// IsActive reports whether this user may authenticate or hold a session.
|
||||
func (u User) IsActive() bool { return u.Status == StatusActive }
|
||||
|
||||
// CanAuthenticate reports whether a password check is even possible. A user
|
||||
// with no password hash cannot sign in, and must be refused in the same way
|
||||
// and with the same timing as a wrong password.
|
||||
func (u User) CanAuthenticate() bool { return u.IsActive() && u.PasswordHash != "" }
|
||||
|
||||
// UserStore is the read side of authentication.
|
||||
//
|
||||
// Deliberately read-only apart from MarkLoggedIn: creating and editing users is
|
||||
// the resource layer's job, and nothing in the sign-in path should be able to
|
||||
// write a role, a status or an organization.
|
||||
type UserStore interface {
|
||||
// FindByEmail resolves the login identifier. Email is citext and globally
|
||||
// unique (migration 000004), so this returns at most one row.
|
||||
FindByEmail(ctx context.Context, email string) (User, error)
|
||||
// FindByID resolves the user behind a session on every request.
|
||||
FindByID(ctx context.Context, id string) (User, error)
|
||||
// MarkLoggedIn records a successful sign-in.
|
||||
MarkLoggedIn(ctx context.Context, id string, at time.Time) error
|
||||
}
|
||||
|
||||
// PGUserStore reads users from PostgreSQL.
|
||||
type PGUserStore struct {
|
||||
db Querier
|
||||
}
|
||||
|
||||
// NewPGUserStore builds the store over a pool or a transaction.
|
||||
func NewPGUserStore(db Querier) *PGUserStore { return &PGUserStore{db: db} }
|
||||
|
||||
var _ UserStore = (*PGUserStore)(nil)
|
||||
|
||||
// userColumns is the projection both lookups share.
|
||||
//
|
||||
// password_hash is COALESCEd to the empty string rather than scanned into a
|
||||
// *string: a NULL hash and an empty hash mean the same thing here — no password
|
||||
// is set — and collapsing them at the edge means no caller has to remember to
|
||||
// nil-check before handing the value to VerifyPassword.
|
||||
const userColumns = `id::text, org_id::text, email::text, full_name,
|
||||
role, account_type, status, COALESCE(password_hash, '')`
|
||||
|
||||
func scanUser(row pgx.Row) (User, error) {
|
||||
var u User
|
||||
err := row.Scan(&u.ID, &u.OrgID, &u.Email, &u.FullName,
|
||||
&u.Role, &u.AccountType, &u.Status, &u.PasswordHash)
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return User{}, ErrUserNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return User{}, err
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// FindByEmail looks a user up by their login identifier.
|
||||
//
|
||||
// The comparison is against the citext column, so it is case-insensitive:
|
||||
// "Demo@Krow.app" finds the same row as "demo@krow.app", which is what a person
|
||||
// typing their own address at a login form expects. Trimming is the caller's
|
||||
// job and is done in the handler, where the raw input is.
|
||||
func (s *PGUserStore) FindByEmail(ctx context.Context, email string) (User, error) {
|
||||
if email == "" {
|
||||
return User{}, ErrUserNotFound
|
||||
}
|
||||
const q = `SELECT ` + userColumns + ` FROM users WHERE email = $1::citext`
|
||||
u, err := scanUser(s.db.QueryRow(ctx, q, email))
|
||||
if err != nil && !errors.Is(err, ErrUserNotFound) {
|
||||
return User{}, fmt.Errorf("auth: find user by email: %w", err)
|
||||
}
|
||||
return u, err
|
||||
}
|
||||
|
||||
// FindByID resolves the user behind a session.
|
||||
//
|
||||
// This runs on every authenticated request, which is why status is read here
|
||||
// rather than cached in the session row: suspending an account must take effect
|
||||
// on the next request, not whenever the session happens to expire.
|
||||
func (s *PGUserStore) FindByID(ctx context.Context, id string) (User, error) {
|
||||
if id == "" {
|
||||
return User{}, ErrUserNotFound
|
||||
}
|
||||
const q = `SELECT ` + userColumns + ` FROM users WHERE id = $1::uuid`
|
||||
u, err := scanUser(s.db.QueryRow(ctx, q, id))
|
||||
if err != nil && !errors.Is(err, ErrUserNotFound) {
|
||||
return User{}, fmt.Errorf("auth: find user by id: %w", err)
|
||||
}
|
||||
return u, err
|
||||
}
|
||||
|
||||
// MarkLoggedIn stamps last_login_at.
|
||||
//
|
||||
// Deliberately not part of the transaction that creates the session: a failure
|
||||
// to record the timestamp is a lost diagnostic, not a reason to refuse a
|
||||
// sign-in that has already succeeded on its merits.
|
||||
func (s *PGUserStore) MarkLoggedIn(ctx context.Context, id string, at time.Time) error {
|
||||
const q = `UPDATE users SET last_login_at = $2::timestamptz WHERE id = $1::uuid`
|
||||
if _, err := s.db.Exec(ctx, q, id, at); err != nil {
|
||||
return fmt.Errorf("auth: mark logged in: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user