first commit

This commit is contained in:
2026-08-24 13:06:29 +05:30
commit 7d12ebef3d
86 changed files with 39996 additions and 0 deletions

View 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
}

View 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
}

View 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] + "$"
}

View 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(&reg); 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
}

View 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())
}

View 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)
}
}
}

View 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
}

View 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) }

View 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)
}
}
}

View 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
}