add users logic
This commit is contained in:
@@ -0,0 +1,67 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
)
|
||||
|
||||
// Claims is the payload of a signed token. Subject (from RegisteredClaims) is
|
||||
// the user id — it becomes the coordinator's jobs.owner_id; Role drives
|
||||
// authorization. Both services verify this token locally with the shared HS256
|
||||
// secret, so no runtime call back to the userservice is ever needed.
|
||||
type Claims struct {
|
||||
Role domain.Role `json:"role"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// Issuer signs and verifies tokens with a shared HS256 secret.
|
||||
type Issuer struct {
|
||||
secret []byte
|
||||
ttl time.Duration
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
// NewIssuer builds an Issuer. now defaults to time.Now when nil; tests inject a
|
||||
// fixed clock to make expiry deterministic.
|
||||
func NewIssuer(secret string, ttl time.Duration, now func() time.Time) Issuer {
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
return Issuer{secret: []byte(secret), ttl: ttl, now: now}
|
||||
}
|
||||
|
||||
// Issue returns a signed token for the user, valid for the configured TTL.
|
||||
func (i Issuer) Issue(userID uuid.UUID, role domain.Role) (string, error) {
|
||||
now := i.now()
|
||||
claims := Claims{
|
||||
Role: role,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: userID.String(),
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
ExpiresAt: jwt.NewNumericDate(now.Add(i.ttl)),
|
||||
},
|
||||
}
|
||||
return jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(i.secret)
|
||||
}
|
||||
|
||||
// Verify checks the signature and expiry and returns the claims. It pins the
|
||||
// algorithm to HMAC, rejecting a token that asks for "none" or an RS256 public
|
||||
// key — the classic algorithm-substitution attack against naive verifiers.
|
||||
func (i Issuer) Verify(token string) (*Claims, error) {
|
||||
var claims Claims
|
||||
_, err := jwt.ParseWithClaims(token, &claims, func(t *jwt.Token) (any, error) {
|
||||
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("unexpected signing method: %v", t.Header["alg"])
|
||||
}
|
||||
return i.secret, nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &claims, nil
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
)
|
||||
|
||||
const testSecret = "test-secret-at-least-32-bytes-long!!"
|
||||
|
||||
func TestIssueVerifyRoundTrip(t *testing.T) {
|
||||
iss := NewIssuer(testSecret, time.Hour, nil)
|
||||
id := uuid.New()
|
||||
|
||||
token, err := iss.Issue(id, domain.RoleAdmin)
|
||||
if err != nil {
|
||||
t.Fatalf("issue: %v", err)
|
||||
}
|
||||
|
||||
claims, err := iss.Verify(token)
|
||||
if err != nil {
|
||||
t.Fatalf("verify: %v", err)
|
||||
}
|
||||
if claims.Subject != id.String() {
|
||||
t.Errorf("sub = %q, want %q", claims.Subject, id.String())
|
||||
}
|
||||
if claims.Role != domain.RoleAdmin {
|
||||
t.Errorf("role = %q, want admin", claims.Role)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyRejectsExpired(t *testing.T) {
|
||||
// Negative TTL: the token is already expired when issued.
|
||||
iss := NewIssuer(testSecret, -time.Minute, nil)
|
||||
token, _ := iss.Issue(uuid.New(), domain.RoleUser)
|
||||
|
||||
if _, err := iss.Verify(token); err == nil {
|
||||
t.Error("expired token accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyRejectsWrongSecret(t *testing.T) {
|
||||
token, _ := NewIssuer(testSecret, time.Hour, nil).Issue(uuid.New(), domain.RoleUser)
|
||||
|
||||
other := NewIssuer("another-secret-also-32-bytes-long!!!", time.Hour, nil)
|
||||
if _, err := other.Verify(token); err == nil {
|
||||
t.Error("token verified under the wrong secret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyRejectsNoneAlgorithm(t *testing.T) {
|
||||
// Forge a token signed with "none" — the classic algorithm-substitution
|
||||
// attack. A verifier that trusts the header's alg would accept it.
|
||||
tok := jwt.NewWithClaims(jwt.SigningMethodNone, Claims{
|
||||
Role: domain.RoleAdmin,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: uuid.New().String(),
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
|
||||
},
|
||||
})
|
||||
raw, err := tok.SignedString(jwt.UnsafeAllowNoneSignatureType)
|
||||
if err != nil {
|
||||
t.Fatalf("sign none: %v", err)
|
||||
}
|
||||
|
||||
iss := NewIssuer(testSecret, time.Hour, nil)
|
||||
if _, err := iss.Verify(raw); err == nil {
|
||||
t.Error("none-signed token accepted")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
// Package auth holds the cryptographic adapters — password hashing and JWT
|
||||
// signing/verification. They implement use-case ports and keep bcrypt and the
|
||||
// JWT library out of the domain and use-case layers.
|
||||
package auth
|
||||
|
||||
import "golang.org/x/crypto/bcrypt"
|
||||
|
||||
// Hasher turns plaintext passwords into storable hashes and checks them back.
|
||||
type Hasher struct {
|
||||
cost int
|
||||
}
|
||||
|
||||
// NewHasher builds a Hasher. A cost of 0 uses bcrypt's default work factor.
|
||||
func NewHasher(cost int) Hasher {
|
||||
if cost == 0 {
|
||||
cost = bcrypt.DefaultCost
|
||||
}
|
||||
return Hasher{cost: cost}
|
||||
}
|
||||
|
||||
// Hash returns the bcrypt hash of password. The salt and the cost are embedded
|
||||
// in the returned string, so nothing else needs to be stored alongside it.
|
||||
//
|
||||
// bcrypt silently ignores input past 72 bytes; the use case rejects longer
|
||||
// passwords before reaching here so a truncated tail never becomes a security
|
||||
// surprise.
|
||||
func (h Hasher) Hash(password string) (string, error) {
|
||||
b, err := bcrypt.GenerateFromPassword([]byte(password), h.cost)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(b), nil
|
||||
}
|
||||
|
||||
// Compare reports whether password matches the stored hash. It returns a
|
||||
// non-nil error (bcrypt.ErrMismatchedHashAndPassword) on any mismatch, which
|
||||
// the caller collapses into a generic authentication failure.
|
||||
func (h Hasher) Compare(hash, password string) error {
|
||||
return bcrypt.CompareHashAndPassword([]byte(hash), []byte(password))
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
package auth
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestHashAndCompare(t *testing.T) {
|
||||
h := NewHasher(0) // default cost
|
||||
|
||||
hash, err := h.Hash("correct horse battery staple")
|
||||
if err != nil {
|
||||
t.Fatalf("hash: %v", err)
|
||||
}
|
||||
if hash == "correct horse battery staple" {
|
||||
t.Fatal("hash must not equal the plaintext")
|
||||
}
|
||||
if err := h.Compare(hash, "correct horse battery staple"); err != nil {
|
||||
t.Errorf("correct password rejected: %v", err)
|
||||
}
|
||||
if err := h.Compare(hash, "wrong password"); err == nil {
|
||||
t.Error("wrong password accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHashSaltsEachTime(t *testing.T) {
|
||||
h := NewHasher(0)
|
||||
a, _ := h.Hash("same")
|
||||
b, _ := h.Hash("same")
|
||||
if a == b {
|
||||
t.Error("two hashes of the same password must differ (random salt)")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package domain
|
||||
|
||||
import "errors"
|
||||
|
||||
// Domain validation errors. They describe an entity that cannot be constructed,
|
||||
// independent of storage or transport, and the HTTP layer maps them to 400.
|
||||
var (
|
||||
ErrEmptyEmail = errors.New("email is required")
|
||||
ErrInvalidEmail = errors.New("email is not a valid address")
|
||||
ErrEmptyPasswordHash = errors.New("password hash is required")
|
||||
)
|
||||
@@ -0,0 +1,80 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"net/mail"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
type Role string
|
||||
|
||||
const (
|
||||
RoleAdmin Role = "admin"
|
||||
RoleUser Role = "user"
|
||||
)
|
||||
|
||||
func (r Role) Valid() bool {
|
||||
switch r {
|
||||
case RoleAdmin, RoleUser:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
type User struct {
|
||||
ID uuid.UUID
|
||||
Email string
|
||||
PasswordHash string
|
||||
Role Role
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
|
||||
// NewUser builds a freshly registered account. It normalises the email and
|
||||
// enforces every invariant a row must satisfy, so an invalid User cannot be
|
||||
// constructed. The caller supplies the already-hashed password — hashing is an
|
||||
// adapter's job, not the domain's.
|
||||
//
|
||||
// Registration always produces a plain user; promotion to admin is a manual,
|
||||
// out-of-band operation, never something a request can trigger.
|
||||
func NewUser(email, passwordHash string, now time.Time) (*User, error) {
|
||||
email = NormalizeEmail(email)
|
||||
if err := validateEmail(email); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if passwordHash == "" {
|
||||
return nil, ErrEmptyPasswordHash
|
||||
}
|
||||
return &User{
|
||||
ID: uuid.New(),
|
||||
Email: email,
|
||||
PasswordHash: passwordHash,
|
||||
Role: RoleUser,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// NormalizeEmail lower-cases and trims an address so that "Bob@X.com " and
|
||||
// "bob@x.com" resolve to the same account. Every lookup and every insert must
|
||||
// pass through here, matching the ck_users_email_lower database constraint.
|
||||
func NormalizeEmail(email string) string {
|
||||
return strings.ToLower(strings.TrimSpace(email))
|
||||
}
|
||||
|
||||
func validateEmail(email string) error {
|
||||
if email == "" {
|
||||
return ErrEmptyEmail
|
||||
}
|
||||
// A minimal shape check, not full RFC 5322: real deliverability is proven by
|
||||
// sending mail, not by a regex. mail.ParseAddress also accepts the
|
||||
// "Name <addr>" form, so we insist the parsed address equals the input.
|
||||
addr, err := mail.ParseAddress(email)
|
||||
if err != nil || addr.Address != email {
|
||||
return ErrInvalidEmail
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
package domain
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
func TestNewUserNormalisesAndValidates(t *testing.T) {
|
||||
now := time.Date(2026, 7, 26, 12, 0, 0, 0, time.UTC)
|
||||
|
||||
u, err := NewUser(" Bob@Example.COM ", "hashed", now)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if u.Email != "bob@example.com" {
|
||||
t.Errorf("email not normalised: got %q", u.Email)
|
||||
}
|
||||
if u.Role != RoleUser {
|
||||
t.Errorf("new user must default to RoleUser, got %q", u.Role)
|
||||
}
|
||||
if u.ID == uuid.Nil {
|
||||
t.Error("new user must get an id")
|
||||
}
|
||||
if !u.CreatedAt.Equal(now) || !u.UpdatedAt.Equal(now) {
|
||||
t.Error("timestamps not set from clock")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewUserRejectsBadInput(t *testing.T) {
|
||||
now := time.Now()
|
||||
cases := []struct {
|
||||
name string
|
||||
email string
|
||||
hash string
|
||||
wantErr error
|
||||
}{
|
||||
{"empty email", "", "h", ErrEmptyEmail},
|
||||
{"no domain", "bob", "h", ErrInvalidEmail},
|
||||
{"name form", "Bob <bob@x.com>", "h", ErrInvalidEmail},
|
||||
{"empty hash", "bob@x.com", "", ErrEmptyPasswordHash},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := NewUser(tc.email, tc.hash, now)
|
||||
if !errors.Is(err, tc.wantErr) {
|
||||
t.Errorf("got %v, want %v", err, tc.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRoleValid(t *testing.T) {
|
||||
if !RoleUser.Valid() || !RoleAdmin.Valid() {
|
||||
t.Error("user and admin must be valid")
|
||||
}
|
||||
if Role("root").Valid() {
|
||||
t.Error("unknown role must be invalid")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
// Clock: the real implementation of the usecase.Clock port. It lives out here
|
||||
// because reading the system clock is infrastructure; tests substitute a fixed one.
|
||||
package infra
|
||||
|
||||
import "time"
|
||||
|
||||
type System struct{}
|
||||
|
||||
func NewClock() System { return System{} }
|
||||
|
||||
// Now returns UTC so every timestamp the coordinator writes is comparable
|
||||
// regardless of the host's timezone.
|
||||
func (System) Now() time.Time { return time.Now().UTC() }
|
||||
@@ -0,0 +1,151 @@
|
||||
// Config: userservice settings, read only from the environment, so the same
|
||||
// binary behaves identically in CI, local, and prod.
|
||||
package infra
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"math"
|
||||
"os"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/joho/godotenv"
|
||||
)
|
||||
|
||||
// defaultEnvFile is loaded by LoadConfig unless ENV_FILE points elsewhere.
|
||||
const defaultEnvFile = ".env"
|
||||
|
||||
type Config struct {
|
||||
// HTTP listen address, e.g. ":8081".
|
||||
Addr string
|
||||
// PostgreSQL connection string (pgx format / libpq URL).
|
||||
DatabaseURL string
|
||||
|
||||
// Shared HS256 secret used to sign JWTs. The coordinator verifies tokens
|
||||
// with this same secret, so the two values MUST match. This is the only
|
||||
// secret shared between the services.
|
||||
JWTSecret string
|
||||
// How long an issued token stays valid.
|
||||
TokenTTL time.Duration
|
||||
// bcrypt work factor. 0 falls back to the library default (currently 10).
|
||||
BcryptCost int
|
||||
|
||||
// Minimum log level: debug, info, warn, error.
|
||||
LogLevel string
|
||||
// Path to a rotated log file. Empty logs to stdout only.
|
||||
LogFile string
|
||||
|
||||
// Connection pool upper bound.
|
||||
DBMaxConns int32
|
||||
// How long to keep retrying the initial database connection at startup
|
||||
// before giving up. Covers a Postgres container that is still booting.
|
||||
DBConnectTimeout time.Duration
|
||||
// Per-request timeout applied to every handler.
|
||||
RequestTimeout time.Duration
|
||||
}
|
||||
|
||||
// LoadConfig reads the environment and fails fast on anything required-but-
|
||||
// missing or malformed, so a misconfigured process never limps along half-wired.
|
||||
//
|
||||
// A .env file (path overridable via ENV_FILE) is loaded first as a local-dev
|
||||
// convenience. It only fills variables the environment does not already define.
|
||||
func LoadConfig() (Config, error) {
|
||||
envFile := os.Getenv("ENV_FILE")
|
||||
if envFile == "" {
|
||||
envFile = defaultEnvFile
|
||||
}
|
||||
// godotenv.Load never overwrites variables already present in the
|
||||
// environment, so an orchestrator's values always beat the file. A missing
|
||||
// file is expected in production, where env vars are injected directly.
|
||||
if err := godotenv.Load(envFile); err != nil && !errors.Is(err, fs.ErrNotExist) {
|
||||
return Config{}, fmt.Errorf("load env file %q: %w", envFile, err)
|
||||
}
|
||||
|
||||
cfg := Config{
|
||||
Addr: getEnv("USERSERVICE_ADDR", ":8081"),
|
||||
DatabaseURL: os.Getenv("DATABASE_URL"),
|
||||
JWTSecret: os.Getenv("JWT_SECRET"),
|
||||
LogLevel: getEnv("LOG_LEVEL", "info"),
|
||||
LogFile: os.Getenv("LOG_FILE"),
|
||||
TokenTTL: 24 * time.Hour,
|
||||
DBMaxConns: 10,
|
||||
DBConnectTimeout: 30 * time.Second,
|
||||
RequestTimeout: 15 * time.Second,
|
||||
}
|
||||
|
||||
if cfg.DatabaseURL == "" {
|
||||
return Config{}, fmt.Errorf("DATABASE_URL is required")
|
||||
}
|
||||
if cfg.JWTSecret == "" {
|
||||
return Config{}, fmt.Errorf("JWT_SECRET is required")
|
||||
}
|
||||
// A short secret makes the HMAC brute-forceable; refuse to start with one.
|
||||
if len(cfg.JWTSecret) < 32 {
|
||||
return Config{}, fmt.Errorf("JWT_SECRET must be at least 32 bytes")
|
||||
}
|
||||
|
||||
var err error
|
||||
if cfg.TokenTTL, err = getEnvDuration("JWT_TTL", cfg.TokenTTL); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.BcryptCost, err = getEnvInt("BCRYPT_COST", cfg.BcryptCost); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.DBMaxConns, err = getEnvInt32("DB_MAX_CONNS", cfg.DBMaxConns); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.DBConnectTimeout, err = getEnvDuration("DB_CONNECT_TIMEOUT", cfg.DBConnectTimeout); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
if cfg.RequestTimeout, err = getEnvDuration("REQUEST_TIMEOUT", cfg.RequestTimeout); err != nil {
|
||||
return Config{}, err
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func getEnv(key, def string) string {
|
||||
if v := os.Getenv(key); v != "" {
|
||||
return v
|
||||
}
|
||||
return def
|
||||
}
|
||||
|
||||
func getEnvInt(key string, def int) (int, error) {
|
||||
v := os.Getenv(key)
|
||||
if v == "" {
|
||||
return def, nil
|
||||
}
|
||||
n, err := strconv.Atoi(v)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("%s: %w", key, err)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func getEnvInt32(key string, def int32) (int32, error) {
|
||||
n, err := getEnvInt(key, int(def))
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
// On 64-bit builds int is wider than int32, so an oversized value would
|
||||
// wrap silently — DB_MAX_CONNS=2147483648 becoming a negative pool size.
|
||||
if n < math.MinInt32 || n > math.MaxInt32 {
|
||||
return 0, fmt.Errorf("%s: %d is out of range for int32", key, n)
|
||||
}
|
||||
return int32(n), nil
|
||||
}
|
||||
|
||||
func getEnvDuration(key string, def time.Duration) (time.Duration, error) {
|
||||
v := os.Getenv(key)
|
||||
if v == "" {
|
||||
return def, nil
|
||||
}
|
||||
d, err := time.ParseDuration(v)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("%s: %w", key, err)
|
||||
}
|
||||
return d, nil
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
// DB: the PostgreSQL connection pool.
|
||||
package infra
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/cenkalti/backoff/v4"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// NewPool builds the single shared pool. The caller owns its lifetime and must
|
||||
// Close() it on shutdown.
|
||||
func NewPool(ctx context.Context, cfg Config, log *slog.Logger) (*pgxpool.Pool, error) {
|
||||
poolCfg, err := pgxpool.ParseConfig(cfg.DatabaseURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
poolCfg.MaxConns = cfg.DBMaxConns
|
||||
|
||||
pool, err := pgxpool.NewWithConfig(ctx, poolCfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// pgxpool.New is lazy, so a ping is needed to actually reach the server.
|
||||
// It is retried because at startup — especially under docker-compose, where
|
||||
// the coordinator can boot before Postgres is accepting connections — a
|
||||
// service should wait for its database rather than crash-loop.
|
||||
if err := pingWithRetry(ctx, pool, cfg.DBConnectTimeout, log); err != nil {
|
||||
pool.Close()
|
||||
return nil, err
|
||||
}
|
||||
return pool, nil
|
||||
}
|
||||
|
||||
// pingWithRetry waits for the database to accept connections, backing off
|
||||
// between attempts until the budget elapses or ctx is cancelled.
|
||||
//
|
||||
// Unlike the transaction retry in storage/postgres, this retries *any* ping
|
||||
// error: at startup a "connection refused" is the expected, retryable state,
|
||||
// not an anomaly.
|
||||
func pingWithRetry(ctx context.Context, pool *pgxpool.Pool, budget time.Duration, log *slog.Logger) error {
|
||||
b := backoff.NewExponentialBackOff()
|
||||
b.InitialInterval = 200 * time.Millisecond
|
||||
b.MaxInterval = 3 * time.Second
|
||||
b.MaxElapsedTime = budget
|
||||
|
||||
attempt := 0
|
||||
return backoff.RetryNotify(
|
||||
func() error {
|
||||
// A bounded per-attempt timeout so one hung dial cannot eat the
|
||||
// whole budget in a single try.
|
||||
pingCtx, cancel := context.WithTimeout(ctx, 3*time.Second)
|
||||
defer cancel()
|
||||
return pool.Ping(pingCtx)
|
||||
},
|
||||
backoff.WithContext(b, ctx),
|
||||
func(err error, next time.Duration) {
|
||||
attempt++
|
||||
log.Warn("database not ready, retrying",
|
||||
"attempt", attempt, "retry_in", next.String(), "err", err)
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package infra
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"gopkg.in/natefinch/lumberjack.v2"
|
||||
)
|
||||
|
||||
// NewLogger builds the process logger.
|
||||
//
|
||||
// It always writes JSON to stdout, so `docker logs` and any 12-factor log
|
||||
// collector keep working. When LogFile is set it *also* writes to a
|
||||
// size-rotated file, so logs survive a container rebuild instead of vanishing
|
||||
// with the previous stdout stream. Rotation is delegated to lumberjack rather
|
||||
// than hand-rolled.
|
||||
//
|
||||
// The returned Closer flushes and closes the file; call it on shutdown.
|
||||
func NewLogger(cfg Config) (*slog.Logger, io.Closer, error) {
|
||||
opts := &slog.HandlerOptions{Level: parseLevel(cfg.LogLevel)}
|
||||
|
||||
var (
|
||||
out io.Writer = os.Stdout
|
||||
closer io.Closer = noopCloser{}
|
||||
)
|
||||
|
||||
if cfg.LogFile != "" {
|
||||
if err := os.MkdirAll(filepath.Dir(cfg.LogFile), 0o750); err != nil {
|
||||
return nil, nil, fmt.Errorf("create log directory: %w", err)
|
||||
}
|
||||
rotator := &lumberjack.Logger{
|
||||
Filename: cfg.LogFile,
|
||||
MaxSize: 50, // megabytes before a rotation
|
||||
MaxBackups: 5, // keep this many rotated files
|
||||
MaxAge: 30, // days
|
||||
Compress: true,
|
||||
}
|
||||
// Tee to both: the console stays live while the file is the durable copy.
|
||||
out = io.MultiWriter(os.Stdout, rotator)
|
||||
closer = rotator
|
||||
}
|
||||
|
||||
return slog.New(slog.NewJSONHandler(out, opts)), closer, nil
|
||||
}
|
||||
|
||||
func parseLevel(s string) slog.Level {
|
||||
switch strings.ToLower(strings.TrimSpace(s)) {
|
||||
case "debug":
|
||||
return slog.LevelDebug
|
||||
case "warn", "warning":
|
||||
return slog.LevelWarn
|
||||
case "error":
|
||||
return slog.LevelError
|
||||
default:
|
||||
return slog.LevelInfo
|
||||
}
|
||||
}
|
||||
|
||||
type noopCloser struct{}
|
||||
|
||||
func (noopCloser) Close() error { return nil }
|
||||
@@ -0,0 +1,74 @@
|
||||
// Server: the HTTP listener and the background lease reaper, both shut down
|
||||
// cleanly on a signal.
|
||||
package infra
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
const shutdownGrace = 15 * time.Second
|
||||
|
||||
// Run serves handler until ctx is cancelled, then drains in-flight requests.
|
||||
func RunServer(ctx context.Context, log *slog.Logger, addr string, handler http.Handler) error {
|
||||
srv := &http.Server{
|
||||
Addr: addr,
|
||||
Handler: handler,
|
||||
ReadHeaderTimeout: 5 * time.Second,
|
||||
}
|
||||
|
||||
// Buffered so this goroutine can exit even when nobody reads the channel
|
||||
// (the ctx.Done branch below) — an unbuffered send would leak it forever.
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
log.Info("userservice listening", "addr", addr)
|
||||
if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
|
||||
errCh <- err
|
||||
}
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-errCh:
|
||||
return err
|
||||
case <-ctx.Done():
|
||||
log.Info("shutdown signal received")
|
||||
}
|
||||
|
||||
// A fresh context: ctx is already cancelled, and reusing it would abort the
|
||||
// very requests we are trying to let finish.
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), shutdownGrace)
|
||||
defer cancel()
|
||||
return srv.Shutdown(shutdownCtx)
|
||||
}
|
||||
|
||||
// RunReaper periodically reclaims tasks whose lease elapsed, so a worker that
|
||||
// died without a heartbeat cannot strand its task in 'leased' forever.
|
||||
// RunPeriodic invokes fn on an interval until ctx is done, logging how many rows
|
||||
// each tick affected. It backs the background reapers (expired leases, offline
|
||||
// workers) — each is a set-based UPDATE that is safe to run repeatedly and
|
||||
// concurrently across coordinators.
|
||||
func RunPeriodic(ctx context.Context, log *slog.Logger, name string, interval time.Duration,
|
||||
fn func(context.Context) (int64, error)) {
|
||||
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
n, err := fn(ctx)
|
||||
if err != nil {
|
||||
log.Debug(name+" skipped", "err", err)
|
||||
continue
|
||||
}
|
||||
if n > 0 {
|
||||
log.Info(name, "count", n)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
// Package memstore provides in-memory implementations of the usecase ports for
|
||||
// fast, deterministic tests that need no database.
|
||||
package memstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
// UserRepo is an in-memory usecase.UserRepository. It stores copies, so callers
|
||||
// mutating a returned user cannot corrupt the store.
|
||||
type UserRepo struct {
|
||||
mu sync.Mutex
|
||||
byID map[uuid.UUID]domain.User
|
||||
byEmail map[string]uuid.UUID
|
||||
}
|
||||
|
||||
func NewUserRepo() *UserRepo {
|
||||
return &UserRepo{
|
||||
byID: make(map[uuid.UUID]domain.User),
|
||||
byEmail: make(map[string]uuid.UUID),
|
||||
}
|
||||
}
|
||||
|
||||
func (r *UserRepo) Insert(_ context.Context, u *domain.User) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if _, ok := r.byEmail[u.Email]; ok {
|
||||
return usecase.ErrEmailExists
|
||||
}
|
||||
r.byID[u.ID] = *u
|
||||
r.byEmail[u.Email] = u.ID
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *UserRepo) GetByEmail(_ context.Context, email string) (*domain.User, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
id, ok := r.byEmail[email]
|
||||
if !ok {
|
||||
return nil, usecase.ErrUserNotFound
|
||||
}
|
||||
u := r.byID[id]
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
func (r *UserRepo) GetByID(_ context.Context, id uuid.UUID) (*domain.User, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
u, ok := r.byID[id]
|
||||
if !ok {
|
||||
return nil, usecase.ErrUserNotFound
|
||||
}
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
// Clock is a fixed usecase.Clock for deterministic tests.
|
||||
type Clock struct{ T time.Time }
|
||||
|
||||
func (c Clock) Now() time.Time { return c.T }
|
||||
@@ -0,0 +1,11 @@
|
||||
package postgres
|
||||
|
||||
import sq "github.com/Masterminds/squirrel"
|
||||
|
||||
// psql is the shared statement builder, fixed to PostgreSQL $N placeholders so
|
||||
// no call site repeats PlaceholderFormat(sq.Dollar).
|
||||
//
|
||||
// Not everything goes through it. Two genuinely set-based statements stay as
|
||||
// raw SQL — claimNext (a FOR UPDATE SKIP LOCKED CTE) and expireLeases (CASE
|
||||
// logic in the SET) — because a builder would obscure them, not clarify them.
|
||||
var psql = sq.StatementBuilder.PlaceholderFormat(sq.Dollar)
|
||||
@@ -0,0 +1,109 @@
|
||||
//go:build integration
|
||||
|
||||
// Integration tests run against a real PostgreSQL instance supplied through
|
||||
// TEST_DATABASE_URL, with the userservice migrations already applied. A real DB
|
||||
// is required because the guarantees under test — the unique-email constraint
|
||||
// mapping to ErrEmailExists, the ck_users_email_lower check — are properties of
|
||||
// Postgres, not of the Go code.
|
||||
//
|
||||
// docker compose up -d
|
||||
// TEST_DATABASE_URL='postgres://scimesh:scimesh@localhost:5432/scimesh?sslmode=disable' \
|
||||
// go test -tags=integration ./internal/storage/postgres/ -v
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
func testPool(t *testing.T) *pgxpool.Pool {
|
||||
t.Helper()
|
||||
url := os.Getenv("TEST_DATABASE_URL")
|
||||
if url == "" {
|
||||
t.Skip("TEST_DATABASE_URL is not set")
|
||||
}
|
||||
pool, err := pgxpool.New(context.Background(), url)
|
||||
if err != nil {
|
||||
t.Fatalf("connect: %v", err)
|
||||
}
|
||||
t.Cleanup(pool.Close)
|
||||
return pool
|
||||
}
|
||||
|
||||
// seedUser inserts a user with a unique email and removes it afterwards, so
|
||||
// tests stay independent of each other and of leftovers from earlier runs.
|
||||
func seedUser(t *testing.T, repo *UserRepo) *domain.User {
|
||||
t.Helper()
|
||||
email := fmt.Sprintf("it-%s@example.com", uuid.NewString())
|
||||
u, err := domain.NewUser(email, "$2a$04$abcdefghijklmnopqrstuv", time.Now().UTC())
|
||||
if err != nil {
|
||||
t.Fatalf("build user: %v", err)
|
||||
}
|
||||
if err := repo.Insert(context.Background(), u); err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_, _ = repo.pool.Exec(context.Background(), "DELETE FROM users WHERE id = $1", u.ID)
|
||||
})
|
||||
return u
|
||||
}
|
||||
|
||||
func TestUserRepoInsertAndGet(t *testing.T) {
|
||||
repo := NewUserRepo(testPool(t))
|
||||
ctx := context.Background()
|
||||
want := seedUser(t, repo)
|
||||
|
||||
byEmail, err := repo.GetByEmail(ctx, want.Email)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByEmail: %v", err)
|
||||
}
|
||||
if byEmail.ID != want.ID || byEmail.Email != want.Email || byEmail.Role != domain.RoleUser {
|
||||
t.Errorf("GetByEmail mismatch: %+v", byEmail)
|
||||
}
|
||||
|
||||
byID, err := repo.GetByID(ctx, want.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("GetByID: %v", err)
|
||||
}
|
||||
if byID.Email != want.Email {
|
||||
t.Errorf("GetByID mismatch: %+v", byID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserRepoDuplicateEmail(t *testing.T) {
|
||||
repo := NewUserRepo(testPool(t))
|
||||
existing := seedUser(t, repo)
|
||||
|
||||
// A second user with the same email must hit the unique constraint and map
|
||||
// to the port's sentinel error.
|
||||
dup, err := domain.NewUser(existing.Email, "$2a$04$abcdefghijklmnopqrstuv", time.Now().UTC())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = repo.Insert(context.Background(), dup)
|
||||
if !errors.Is(err, usecase.ErrEmailExists) {
|
||||
t.Errorf("got %v, want ErrEmailExists", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserRepoNotFound(t *testing.T) {
|
||||
repo := NewUserRepo(testPool(t))
|
||||
ctx := context.Background()
|
||||
|
||||
if _, err := repo.GetByID(ctx, uuid.New()); !errors.Is(err, usecase.ErrUserNotFound) {
|
||||
t.Errorf("GetByID unknown: got %v, want ErrUserNotFound", err)
|
||||
}
|
||||
if _, err := repo.GetByEmail(ctx, "ghost@example.com"); !errors.Is(err, usecase.ErrUserNotFound) {
|
||||
t.Errorf("GetByEmail unknown: got %v, want ErrUserNotFound", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/cenkalti/backoff/v4"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
)
|
||||
|
||||
// Transient PostgreSQL failures. Under concurrent claiming these are expected
|
||||
// rather than exceptional: two coordinators touching neighbouring rows can
|
||||
// deadlock or fail to serialize, and the correct response is to try again.
|
||||
const (
|
||||
codeSerializationFailure = "40001"
|
||||
codeDeadlockDetected = "40P01"
|
||||
codeTooManyConnections = "53300"
|
||||
codeCannotConnectNow = "57P03"
|
||||
)
|
||||
|
||||
// Retry budget: short and bounded. A worker polling for tasks would rather get
|
||||
// a fast error and poll again than have its request hang for half a minute.
|
||||
const (
|
||||
retryInitialInterval = 50 * time.Millisecond
|
||||
retryMaxInterval = 1 * time.Second
|
||||
retryMaxElapsedTime = 5 * time.Second
|
||||
)
|
||||
|
||||
// isTransient reports whether err is worth retrying.
|
||||
//
|
||||
// The default is *not* to retry: a constraint violation or a syntax error will
|
||||
// fail identically every time, and retrying it only multiplies the damage.
|
||||
func isTransient(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
// A cancelled caller does not want another attempt.
|
||||
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
|
||||
return false
|
||||
}
|
||||
|
||||
var pgErr *pgconn.PgError
|
||||
if errors.As(err, &pgErr) {
|
||||
switch pgErr.Code {
|
||||
case codeSerializationFailure, codeDeadlockDetected,
|
||||
codeTooManyConnections, codeCannotConnectNow:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// Connection-level trouble (dropped socket, closed pool). pgconn knows
|
||||
// whether the query could have been executed before the failure — retrying
|
||||
// a maybe-executed write would risk duplicating it.
|
||||
return pgconn.SafeToRetry(err)
|
||||
}
|
||||
|
||||
// withRetry runs op, retrying only transient database failures with
|
||||
// exponential backoff and jitter, and giving up as soon as ctx is done.
|
||||
//
|
||||
// Jitter matters here: without it, several coordinators that collide once will
|
||||
// retry in lockstep and collide again at exactly the same moment.
|
||||
func withRetry(ctx context.Context, op func(context.Context) error) error {
|
||||
b := backoff.NewExponentialBackOff()
|
||||
b.InitialInterval = retryInitialInterval
|
||||
b.MaxInterval = retryMaxInterval
|
||||
b.MaxElapsedTime = retryMaxElapsedTime
|
||||
// RandomizationFactor defaults to 0.5, which is the jitter.
|
||||
|
||||
return backoff.Retry(func() error {
|
||||
err := op(ctx)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if !isTransient(err) {
|
||||
return backoff.Permanent(err) // stop now, do not burn the budget
|
||||
}
|
||||
return err
|
||||
}, backoff.WithContext(b, ctx))
|
||||
}
|
||||
@@ -0,0 +1,98 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
)
|
||||
|
||||
func TestIsTransient(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
want bool
|
||||
}{
|
||||
{"nil", nil, false},
|
||||
{"serialization failure", &pgconn.PgError{Code: codeSerializationFailure}, true},
|
||||
{"deadlock", &pgconn.PgError{Code: codeDeadlockDetected}, true},
|
||||
{"too many connections", &pgconn.PgError{Code: codeTooManyConnections}, true},
|
||||
// A unique-violation repeats identically forever — retrying is pointless.
|
||||
{"unique violation", &pgconn.PgError{Code: "23505"}, false},
|
||||
{"syntax error", &pgconn.PgError{Code: "42601"}, false},
|
||||
{"context cancelled", context.Canceled, false},
|
||||
{"deadline exceeded", context.DeadlineExceeded, false},
|
||||
{"unknown error", errors.New("boom"), false},
|
||||
// Wrapping must not hide the cause: errors.As walks the chain.
|
||||
{"wrapped deadlock", errors2Wrap(&pgconn.PgError{Code: codeDeadlockDetected}), true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isTransient(tt.err); got != tt.want {
|
||||
t.Errorf("isTransient(%v) = %v, want %v", tt.err, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func errors2Wrap(err error) error {
|
||||
return errors.Join(errors.New("query failed"), err)
|
||||
}
|
||||
|
||||
func TestWithRetrySucceedsAfterTransientFailures(t *testing.T) {
|
||||
calls := 0
|
||||
err := withRetry(context.Background(), func(context.Context) error {
|
||||
calls++
|
||||
if calls < 3 {
|
||||
return &pgconn.PgError{Code: codeSerializationFailure}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if calls != 3 {
|
||||
t.Errorf("calls = %d, want 3", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithRetryStopsOnPermanentError(t *testing.T) {
|
||||
permanent := &pgconn.PgError{Code: "23505"} // unique violation
|
||||
calls := 0
|
||||
|
||||
err := withRetry(context.Background(), func(context.Context) error {
|
||||
calls++
|
||||
return permanent
|
||||
})
|
||||
|
||||
if !errors.Is(err, permanent) {
|
||||
t.Errorf("err = %v, want the original error", err)
|
||||
}
|
||||
if calls != 1 {
|
||||
t.Errorf("calls = %d, want 1 — a permanent error must not be retried", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithRetryHonoursContextCancellation(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
calls := 0
|
||||
start := time.Now()
|
||||
err := withRetry(ctx, func(context.Context) error {
|
||||
calls++
|
||||
return &pgconn.PgError{Code: codeDeadlockDetected}
|
||||
})
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected an error once the context expired")
|
||||
}
|
||||
// Must abort at the deadline, not run the full 5s retry budget.
|
||||
if elapsed := time.Since(start); elapsed > time.Second {
|
||||
t.Errorf("took %v, expected to stop at the context deadline", elapsed)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
// Package postgres implements the usecase repository ports on PostgreSQL.
|
||||
// SQL and pgx types never escape this package.
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
)
|
||||
|
||||
// querier is satisfied by both *pgxpool.Pool and pgx.Tx, letting every
|
||||
// repository method run identically inside or outside a transaction.
|
||||
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)
|
||||
SendBatch(ctx context.Context, b *pgx.Batch) pgx.BatchResults
|
||||
}
|
||||
|
||||
// txKey is an unexported struct type, so no other package can collide with it
|
||||
// or reach the transaction we stash in the context.
|
||||
type txKey struct{}
|
||||
|
||||
// TxManager implements usecase.TxManager.
|
||||
type TxManager struct {
|
||||
pool *pgxpool.Pool
|
||||
}
|
||||
|
||||
func NewTxManager(pool *pgxpool.Pool) *TxManager {
|
||||
return &TxManager{pool: pool}
|
||||
}
|
||||
|
||||
// WithinTx runs fn inside one transaction, committing on success and rolling
|
||||
// back on any error or panic.
|
||||
//
|
||||
// The transaction travels in the context rather than in fn's signature, which
|
||||
// is what lets the usecase layer express "do these repository calls atomically"
|
||||
// without its port ever mentioning pgx.
|
||||
// Retrying happens here, around the whole transaction, and deliberately not
|
||||
// inside the repositories. Once Postgres aborts a transaction with a
|
||||
// serialization failure or deadlock, every further statement in it fails too —
|
||||
// replaying a single query would accomplish nothing. The unit of retry is
|
||||
// Begin → fn → Commit.
|
||||
//
|
||||
// This is safe because fn re-reads its rows (via GetForUpdate) on each attempt,
|
||||
// so a retry starts from the current state rather than stale entities.
|
||||
func (m *TxManager) WithinTx(ctx context.Context, fn func(ctx context.Context) error) error {
|
||||
if _, ok := ctx.Value(txKey{}).(pgx.Tx); ok {
|
||||
// Already inside a transaction — join it. Retrying here would be wrong
|
||||
// twice over: the outer transaction owns the retry, and re-running fn
|
||||
// alone cannot undo what the outer one already wrote.
|
||||
return fn(ctx)
|
||||
}
|
||||
|
||||
return withRetry(ctx, func(ctx context.Context) error {
|
||||
return m.runTx(ctx, fn)
|
||||
})
|
||||
}
|
||||
|
||||
func (m *TxManager) runTx(ctx context.Context, fn func(ctx context.Context) error) error {
|
||||
tx, err := m.pool.Begin(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Rollback after a successful Commit is a no-op, so this defer is safe and
|
||||
// also covers the panic path.
|
||||
defer func() { _ = tx.Rollback(ctx) }()
|
||||
|
||||
if err := fn(context.WithValue(ctx, txKey{}, tx)); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Commit(ctx)
|
||||
}
|
||||
|
||||
// conn returns the transaction bound to ctx, or the pool when there is none.
|
||||
func conn(ctx context.Context, pool *pgxpool.Pool) querier {
|
||||
if tx, ok := ctx.Value(txKey{}).(pgx.Tx); ok {
|
||||
return tx
|
||||
}
|
||||
return pool
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
sq "github.com/Masterminds/squirrel"
|
||||
"github.com/google/uuid"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
// uniqueViolation is PostgreSQL's SQLSTATE for a unique-constraint breach.
|
||||
const uniqueViolation = "23505"
|
||||
|
||||
var userColumns = []string{"id", "email", "password_hash", "role", "created_at", "updated_at"}
|
||||
|
||||
// UserRepo implements usecase.UserRepository on PostgreSQL.
|
||||
type UserRepo struct {
|
||||
pool *pgxpool.Pool
|
||||
}
|
||||
|
||||
func NewUserRepo(pool *pgxpool.Pool) *UserRepo {
|
||||
return &UserRepo{pool: pool}
|
||||
}
|
||||
|
||||
func (r *UserRepo) Insert(ctx context.Context, u *domain.User) error {
|
||||
sql, args, err := psql.Insert("users").
|
||||
Columns(userColumns...).
|
||||
Values(u.ID, u.Email, u.PasswordHash, string(u.Role), u.CreatedAt, u.UpdatedAt).
|
||||
ToSql()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := conn(ctx, r.pool).Exec(ctx, sql, args...); err != nil {
|
||||
// A concurrent insert of the same email surfaces as a unique violation
|
||||
// on uq_users_email; translate it to the port's sentinel so the use
|
||||
// case never sees a driver type.
|
||||
if isUniqueViolation(err) {
|
||||
return usecase.ErrEmailExists
|
||||
}
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *UserRepo) GetByEmail(ctx context.Context, email string) (*domain.User, error) {
|
||||
return r.getBy(ctx, sq.Eq{"email": email})
|
||||
}
|
||||
|
||||
func (r *UserRepo) GetByID(ctx context.Context, id uuid.UUID) (*domain.User, error) {
|
||||
return r.getBy(ctx, sq.Eq{"id": id})
|
||||
}
|
||||
|
||||
func (r *UserRepo) getBy(ctx context.Context, pred sq.Sqlizer) (*domain.User, error) {
|
||||
sql, args, err := psql.Select(userColumns...).From("users").Where(pred).ToSql()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return scanUser(conn(ctx, r.pool).QueryRow(ctx, sql, args...))
|
||||
}
|
||||
|
||||
func scanUser(row pgx.Row) (*domain.User, error) {
|
||||
var (
|
||||
u domain.User
|
||||
role string
|
||||
)
|
||||
if err := row.Scan(&u.ID, &u.Email, &u.PasswordHash, &role, &u.CreatedAt, &u.UpdatedAt); err != nil {
|
||||
if errors.Is(err, pgx.ErrNoRows) {
|
||||
return nil, usecase.ErrUserNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
u.Role = domain.Role(role)
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
func isUniqueViolation(err error) bool {
|
||||
var pgErr *pgconn.PgError
|
||||
return errors.As(err, &pgErr) && pgErr.Code == uniqueViolation
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package http
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
)
|
||||
|
||||
// registerRequest / loginRequest are the JSON bodies clients POST. Kept separate
|
||||
// from the domain so the wire format can evolve without touching the entity.
|
||||
type registerRequest struct {
|
||||
Email string `json:"email"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
type loginRequest struct {
|
||||
Email string `json:"email"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
|
||||
// userResponse is the public view of a user. It never carries the password hash.
|
||||
type userResponse struct {
|
||||
ID string `json:"id"`
|
||||
Email string `json:"email"`
|
||||
Role string `json:"role"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
}
|
||||
|
||||
type loginResponse struct {
|
||||
Token string `json:"token"`
|
||||
User userResponse `json:"user"`
|
||||
}
|
||||
|
||||
func toUserResponse(u *domain.User) userResponse {
|
||||
return userResponse{
|
||||
ID: u.ID.String(),
|
||||
Email: u.Email,
|
||||
Role: string(u.Role),
|
||||
CreatedAt: u.CreatedAt.UTC().Format(time.RFC3339),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package http
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
// maxJSONBody caps a request body. Credentials are tiny; anything larger is a
|
||||
// mistake or an attack, so reject it before allocating.
|
||||
const maxJSONBody = 1 << 20 // 1 MiB
|
||||
|
||||
type errorResponse struct {
|
||||
Error string `json:"error"`
|
||||
RequestID string `json:"request_id,omitempty"`
|
||||
}
|
||||
|
||||
func writeJSON(w http.ResponseWriter, status int, v any) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
_ = json.NewEncoder(w).Encode(v)
|
||||
}
|
||||
|
||||
// writeError maps a domain or use-case error to an HTTP status and a safe
|
||||
// message, logging only genuine server faults (5xx). Client errors (4xx) are
|
||||
// expected and stay out of the error log.
|
||||
func writeError(w http.ResponseWriter, r *http.Request, log *slog.Logger, err error) {
|
||||
status, msg := statusForError(err)
|
||||
if status >= http.StatusInternalServerError {
|
||||
log.Error("request failed",
|
||||
"err", err,
|
||||
"request_id", requestIDFrom(r.Context()),
|
||||
"path", r.URL.Path,
|
||||
)
|
||||
}
|
||||
writeJSON(w, status, errorResponse{Error: msg, RequestID: requestIDFrom(r.Context())})
|
||||
}
|
||||
|
||||
func statusForError(err error) (int, string) {
|
||||
switch {
|
||||
case errors.Is(err, usecase.ErrEmailExists):
|
||||
return http.StatusConflict, "email already registered"
|
||||
case errors.Is(err, usecase.ErrInvalidCredentials):
|
||||
return http.StatusUnauthorized, "invalid email or password"
|
||||
case errors.Is(err, usecase.ErrUserNotFound):
|
||||
return http.StatusNotFound, "user not found"
|
||||
case errors.Is(err, usecase.ErrPasswordTooShort):
|
||||
return http.StatusBadRequest, "password must be at least 8 characters"
|
||||
case errors.Is(err, usecase.ErrPasswordTooLong):
|
||||
return http.StatusBadRequest, "password must be at most 72 bytes"
|
||||
case errors.Is(err, domain.ErrEmptyEmail), errors.Is(err, domain.ErrInvalidEmail):
|
||||
return http.StatusBadRequest, "email is not a valid address"
|
||||
default:
|
||||
// Don't leak internals; the real error is in the log under request_id.
|
||||
return http.StatusInternalServerError, "internal error"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package http
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
// Handlers holds the use cases each endpoint drives.
|
||||
type Handlers struct {
|
||||
register *usecase.Register
|
||||
login *usecase.Login
|
||||
users usecase.UserRepository
|
||||
log *slog.Logger
|
||||
}
|
||||
|
||||
// handleHealth is an unauthenticated liveness probe for the container and load
|
||||
// balancer.
|
||||
func (h *Handlers) handleHealth(w http.ResponseWriter, _ *http.Request) {
|
||||
writeJSON(w, http.StatusOK, map[string]string{"status": "ok"})
|
||||
}
|
||||
|
||||
// handleRegister creates an account. It returns 201 with the public user view,
|
||||
// 409 if the email is taken, or 400 on a malformed body / weak password.
|
||||
func (h *Handlers) handleRegister(w http.ResponseWriter, r *http.Request) {
|
||||
var req registerRequest
|
||||
if !decodeJSON(w, r, &req) {
|
||||
return
|
||||
}
|
||||
u, err := h.register.Execute(r.Context(), req.Email, req.Password)
|
||||
if err != nil {
|
||||
writeError(w, r, h.log, err)
|
||||
return
|
||||
}
|
||||
writeJSON(w, http.StatusCreated, toUserResponse(u))
|
||||
}
|
||||
|
||||
// handleLogin verifies credentials and returns a signed token plus the user.
|
||||
func (h *Handlers) handleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
var req loginRequest
|
||||
if !decodeJSON(w, r, &req) {
|
||||
return
|
||||
}
|
||||
token, u, err := h.login.Execute(r.Context(), req.Email, req.Password)
|
||||
if err != nil {
|
||||
writeError(w, r, h.log, err)
|
||||
return
|
||||
}
|
||||
writeJSON(w, http.StatusOK, loginResponse{Token: token, User: toUserResponse(u)})
|
||||
}
|
||||
|
||||
// handleMe returns the caller's own account, proving the token works end to end.
|
||||
// It reads the user id the JWT middleware stashed in the context.
|
||||
func (h *Handlers) handleMe(w http.ResponseWriter, r *http.Request) {
|
||||
id, ok := userIDFrom(r.Context())
|
||||
if !ok {
|
||||
unauthorized(w, r)
|
||||
return
|
||||
}
|
||||
u, err := h.users.GetByID(r.Context(), id)
|
||||
if err != nil {
|
||||
writeError(w, r, h.log, err)
|
||||
return
|
||||
}
|
||||
writeJSON(w, http.StatusOK, toUserResponse(u))
|
||||
}
|
||||
|
||||
// decodeJSON reads a size-capped JSON body into dst, rejecting unknown fields.
|
||||
// It writes a 400 and returns false on any problem, so callers can `if
|
||||
// !decodeJSON(...) { return }`.
|
||||
func decodeJSON(w http.ResponseWriter, r *http.Request, dst any) bool {
|
||||
r.Body = http.MaxBytesReader(w, r.Body, maxJSONBody)
|
||||
dec := json.NewDecoder(r.Body)
|
||||
dec.DisallowUnknownFields()
|
||||
if err := dec.Decode(dst); err != nil {
|
||||
writeJSON(w, http.StatusBadRequest, errorResponse{
|
||||
Error: "invalid JSON body",
|
||||
RequestID: requestIDFrom(r.Context()),
|
||||
})
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package http
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/auth"
|
||||
)
|
||||
|
||||
type ctxKey string
|
||||
|
||||
const (
|
||||
requestIDKey ctxKey = "request_id"
|
||||
userIDKey ctxKey = "user_id"
|
||||
roleKey ctxKey = "role"
|
||||
)
|
||||
|
||||
// withRequestID stamps every request with an ID for correlated logs and error
|
||||
// bodies. It wraps the auth middleware rather than the other way round, so even
|
||||
// a rejected request carries an ID the caller can quote in a bug report.
|
||||
func withRequestID(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
id := newRequestID()
|
||||
w.Header().Set("X-Request-ID", id)
|
||||
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), requestIDKey, id)))
|
||||
})
|
||||
}
|
||||
|
||||
func requestIDFrom(ctx context.Context) string {
|
||||
if v, ok := ctx.Value(requestIDKey).(string); ok {
|
||||
return v
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func newRequestID() string {
|
||||
var b [8]byte
|
||||
_, _ = rand.Read(b[:])
|
||||
return hex.EncodeToString(b[:])
|
||||
}
|
||||
|
||||
// tokenVerifier is the slice of auth.Issuer the JWT middleware needs. Taking an
|
||||
// interface keeps the middleware testable with a stub verifier.
|
||||
type tokenVerifier interface {
|
||||
Verify(token string) (*auth.Claims, error)
|
||||
}
|
||||
|
||||
// withJWT verifies the Bearer token and stashes the caller's id and role in the
|
||||
// request context. It rejects any request without a valid, unexpired HS256
|
||||
// token — this is what protects endpoints that act on a specific user.
|
||||
func withJWT(v tokenVerifier) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
raw := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ")
|
||||
if raw == "" {
|
||||
unauthorized(w, r)
|
||||
return
|
||||
}
|
||||
claims, err := v.Verify(raw)
|
||||
if err != nil {
|
||||
unauthorized(w, r)
|
||||
return
|
||||
}
|
||||
id, err := uuid.Parse(claims.Subject)
|
||||
if err != nil {
|
||||
unauthorized(w, r)
|
||||
return
|
||||
}
|
||||
ctx := context.WithValue(r.Context(), userIDKey, id)
|
||||
ctx = context.WithValue(ctx, roleKey, claims.Role)
|
||||
next.ServeHTTP(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func unauthorized(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("WWW-Authenticate", "Bearer")
|
||||
writeJSON(w, http.StatusUnauthorized, errorResponse{
|
||||
Error: "unauthorized",
|
||||
RequestID: requestIDFrom(r.Context()),
|
||||
})
|
||||
}
|
||||
|
||||
// userIDFrom returns the authenticated caller's id, set by withJWT.
|
||||
func userIDFrom(ctx context.Context) (uuid.UUID, bool) {
|
||||
id, ok := ctx.Value(userIDKey).(uuid.UUID)
|
||||
return id, ok
|
||||
}
|
||||
|
||||
// statusRecorder captures the status code for the access log.
|
||||
type statusRecorder struct {
|
||||
http.ResponseWriter
|
||||
status int
|
||||
}
|
||||
|
||||
func (s *statusRecorder) WriteHeader(code int) {
|
||||
s.status = code
|
||||
s.ResponseWriter.WriteHeader(code)
|
||||
}
|
||||
|
||||
// withAccessLog records one structured line per request — the minimum needed to
|
||||
// debug a distributed system after the fact.
|
||||
func withAccessLog(log *slog.Logger) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
start := time.Now()
|
||||
rec := &statusRecorder{ResponseWriter: w, status: http.StatusOK}
|
||||
next.ServeHTTP(rec, r)
|
||||
log.Info("request",
|
||||
"request_id", requestIDFrom(r.Context()),
|
||||
"method", r.Method,
|
||||
"path", r.URL.Path,
|
||||
"status", rec.status,
|
||||
"duration_ms", time.Since(start).Milliseconds(),
|
||||
)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// chain applies middleware so that the first argument is the outermost layer.
|
||||
func chain(h http.Handler, mw ...func(http.Handler) http.Handler) http.Handler {
|
||||
for i := len(mw) - 1; i >= 0; i-- {
|
||||
h = mw[i](h)
|
||||
}
|
||||
return h
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
// Package http exposes the userservice over HTTP: registration, login, and a
|
||||
// token-protected /me. It owns routing, request decoding, and error mapping;
|
||||
// business rules live in the usecase layer.
|
||||
package http
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/auth"
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
// UseCases bundles the application services the handlers drive.
|
||||
type UseCases struct {
|
||||
Register *usecase.Register
|
||||
Login *usecase.Login
|
||||
Users usecase.UserRepository
|
||||
}
|
||||
|
||||
// NewServer wires the routes and the middleware stack and returns the handler.
|
||||
// The issuer verifies tokens for the protected /me route.
|
||||
func NewServer(log *slog.Logger, uc UseCases, issuer auth.Issuer) http.Handler {
|
||||
h := &Handlers{
|
||||
register: uc.Register,
|
||||
login: uc.Login,
|
||||
users: uc.Users,
|
||||
log: log,
|
||||
}
|
||||
|
||||
mux := http.NewServeMux()
|
||||
// Method-aware patterns (Go 1.22+): a GET to /register is a 405, not a match.
|
||||
mux.HandleFunc("GET /health", h.handleHealth)
|
||||
mux.HandleFunc("POST /register", h.handleRegister)
|
||||
mux.HandleFunc("POST /login", h.handleLogin)
|
||||
// /me proves a token round-trips; it sits behind JWT auth.
|
||||
mux.Handle("GET /me", chain(http.HandlerFunc(h.handleMe), withJWT(issuer)))
|
||||
|
||||
// Outermost first: every request gets an ID and an access-log line.
|
||||
return chain(mux, withRequestID, withAccessLog(log))
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
package http_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/auth"
|
||||
"github.com/emil28092005/SciMesh/users/internal/memstore"
|
||||
apihttp "github.com/emil28092005/SciMesh/users/internal/transport/http"
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
const secret = "server-test-secret-32-bytes-long!!!!"
|
||||
|
||||
func newTestServer() http.Handler {
|
||||
users := memstore.NewUserRepo()
|
||||
hasher := auth.NewHasher(4)
|
||||
clk := memstore.Clock{T: time.Date(2026, 7, 26, 0, 0, 0, 0, time.UTC)}
|
||||
// Real clock for the issuer so tokens are valid at verification time.
|
||||
issuer := auth.NewIssuer(secret, time.Hour, nil)
|
||||
|
||||
uc := apihttp.UseCases{
|
||||
Register: usecase.NewRegister(users, hasher, clk),
|
||||
Login: usecase.NewLogin(users, hasher, issuer),
|
||||
Users: users,
|
||||
}
|
||||
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
return apihttp.NewServer(log, uc, issuer)
|
||||
}
|
||||
|
||||
func do(t *testing.T, h http.Handler, method, path, token string, body any) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
var buf bytes.Buffer
|
||||
if body != nil {
|
||||
if err := json.NewEncoder(&buf).Encode(body); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
req := httptest.NewRequest(method, path, &buf)
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
h.ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
func TestRegisterThenLoginThenMe(t *testing.T) {
|
||||
h := newTestServer()
|
||||
creds := map[string]string{"email": "flow@example.com", "password": "password123"}
|
||||
|
||||
// Register -> 201
|
||||
rec := do(t, h, http.MethodPost, "/register", "", creds)
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("register: got %d, body %s", rec.Code, rec.Body)
|
||||
}
|
||||
|
||||
// Login -> 200 with a token
|
||||
rec = do(t, h, http.MethodPost, "/login", "", creds)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("login: got %d, body %s", rec.Code, rec.Body)
|
||||
}
|
||||
var lr struct {
|
||||
Token string `json:"token"`
|
||||
User struct {
|
||||
Email string `json:"email"`
|
||||
Role string `json:"role"`
|
||||
} `json:"user"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &lr); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if lr.Token == "" || lr.User.Email != "flow@example.com" || lr.User.Role != "user" {
|
||||
t.Fatalf("unexpected login body: %+v", lr)
|
||||
}
|
||||
|
||||
// /me with the token -> 200, same user
|
||||
rec = do(t, h, http.MethodGet, "/me", lr.Token, nil)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("me: got %d, body %s", rec.Code, rec.Body)
|
||||
}
|
||||
var me struct {
|
||||
Email string `json:"email"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &me); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if me.Email != "flow@example.com" {
|
||||
t.Errorf("me email = %q", me.Email)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterDuplicate(t *testing.T) {
|
||||
h := newTestServer()
|
||||
creds := map[string]string{"email": "dup@example.com", "password": "password123"}
|
||||
_ = do(t, h, http.MethodPost, "/register", "", creds)
|
||||
|
||||
rec := do(t, h, http.MethodPost, "/register", "", creds)
|
||||
if rec.Code != http.StatusConflict {
|
||||
t.Errorf("duplicate register: got %d, want 409", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterValidation(t *testing.T) {
|
||||
h := newTestServer()
|
||||
cases := []struct {
|
||||
name string
|
||||
body map[string]string
|
||||
}{
|
||||
{"weak password", map[string]string{"email": "a@b.com", "password": "short"}},
|
||||
{"bad email", map[string]string{"email": "nope", "password": "password123"}},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
rec := do(t, h, http.MethodPost, "/register", "", tc.body)
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("got %d, want 400", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterRejectsUnknownFields(t *testing.T) {
|
||||
h := newTestServer()
|
||||
rec := do(t, h, http.MethodPost, "/register", "", map[string]string{
|
||||
"email": "a@b.com", "password": "password123", "role": "admin",
|
||||
})
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("unknown field must be rejected: got %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginWrongPassword(t *testing.T) {
|
||||
h := newTestServer()
|
||||
_ = do(t, h, http.MethodPost, "/register", "", map[string]string{
|
||||
"email": "x@example.com", "password": "password123",
|
||||
})
|
||||
rec := do(t, h, http.MethodPost, "/login", "", map[string]string{
|
||||
"email": "x@example.com", "password": "wrongpass1",
|
||||
})
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("got %d, want 401", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMeRequiresToken(t *testing.T) {
|
||||
h := newTestServer()
|
||||
if rec := do(t, h, http.MethodGet, "/me", "", nil); rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("no token: got %d, want 401", rec.Code)
|
||||
}
|
||||
if rec := do(t, h, http.MethodGet, "/me", "garbage.token.here", nil); rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("bad token: got %d, want 401", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHealth(t *testing.T) {
|
||||
h := newTestServer()
|
||||
if rec := do(t, h, http.MethodGet, "/health", "", nil); rec.Code != http.StatusOK {
|
||||
t.Errorf("health: got %d", rec.Code)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package usecase
|
||||
|
||||
import "errors"
|
||||
|
||||
var (
|
||||
// Repository-contract errors, returned by UserRepository implementations.
|
||||
ErrEmailExists = errors.New("email already registered")
|
||||
ErrUserNotFound = errors.New("user not found")
|
||||
|
||||
// Use-case errors surfaced to the transport layer.
|
||||
//
|
||||
// ErrInvalidCredentials is deliberately returned for both an unknown email
|
||||
// and a wrong password, so an attacker cannot use the response to learn
|
||||
// which emails are registered.
|
||||
ErrInvalidCredentials = errors.New("invalid email or password")
|
||||
ErrPasswordTooShort = errors.New("password too short")
|
||||
ErrPasswordTooLong = errors.New("password too long")
|
||||
)
|
||||
@@ -0,0 +1,42 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
)
|
||||
|
||||
// Login verifies credentials and issues a signed token.
|
||||
type Login struct {
|
||||
users UserRepository
|
||||
hasher PasswordHasher
|
||||
tokens TokenIssuer
|
||||
}
|
||||
|
||||
func NewLogin(users UserRepository, hasher PasswordHasher, tokens TokenIssuer) *Login {
|
||||
return &Login{users: users, hasher: hasher, tokens: tokens}
|
||||
}
|
||||
|
||||
// Execute returns a signed token and the user on success. It returns
|
||||
// ErrInvalidCredentials for both an unknown email and a wrong password so the
|
||||
// two cases are indistinguishable to a caller probing for valid accounts.
|
||||
func (l *Login) Execute(ctx context.Context, email, password string) (string, *domain.User, error) {
|
||||
u, err := l.users.GetByEmail(ctx, domain.NormalizeEmail(email))
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrUserNotFound) {
|
||||
return "", nil, ErrInvalidCredentials
|
||||
}
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
if err := l.hasher.Compare(u.PasswordHash, password); err != nil {
|
||||
return "", nil, ErrInvalidCredentials
|
||||
}
|
||||
|
||||
token, err := l.tokens.Issue(u.ID, u.Role)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
return token, u, nil
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
// Package usecase holds the application logic — registration and login — plus
|
||||
// the ports (interfaces) it depends on. The concrete adapters (PostgreSQL,
|
||||
// bcrypt, JWT) are injected from cmd, so this package never imports them.
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
)
|
||||
|
||||
// UserRepository persists and looks up users. Implementations return the
|
||||
// sentinel errors in errors.go so the use cases can react without knowing about
|
||||
// SQL or driver types.
|
||||
type UserRepository interface {
|
||||
// Insert stores a new user, returning ErrEmailExists if the email is taken.
|
||||
Insert(ctx context.Context, u *domain.User) error
|
||||
// GetByEmail returns the user with the (normalised) email, or ErrUserNotFound.
|
||||
GetByEmail(ctx context.Context, email string) (*domain.User, error)
|
||||
// GetByID returns the user with id, or ErrUserNotFound.
|
||||
GetByID(ctx context.Context, id uuid.UUID) (*domain.User, error)
|
||||
}
|
||||
|
||||
// PasswordHasher hashes and verifies passwords. The bcrypt adapter satisfies it.
|
||||
type PasswordHasher interface {
|
||||
Hash(password string) (string, error)
|
||||
Compare(hash, password string) error
|
||||
}
|
||||
|
||||
// TokenIssuer mints a signed access token for an authenticated user.
|
||||
type TokenIssuer interface {
|
||||
Issue(userID uuid.UUID, role domain.Role) (string, error)
|
||||
}
|
||||
|
||||
// Clock reads the current time; a fake one makes tests deterministic.
|
||||
type Clock interface {
|
||||
Now() time.Time
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
package usecase
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
)
|
||||
|
||||
const (
|
||||
// minPasswordLen is a floor, not a policy engine — enough to reject the
|
||||
// obviously weak without pretending to measure real strength.
|
||||
minPasswordLen = 8
|
||||
// maxPasswordLen is bcrypt's hard input limit: it ignores bytes past 72, so
|
||||
// accepting a longer password would silently hash only its prefix.
|
||||
maxPasswordLen = 72
|
||||
)
|
||||
|
||||
// Register creates a new account: it validates the password, hashes it, builds
|
||||
// the domain user, and persists it.
|
||||
type Register struct {
|
||||
users UserRepository
|
||||
hasher PasswordHasher
|
||||
clk Clock
|
||||
}
|
||||
|
||||
func NewRegister(users UserRepository, hasher PasswordHasher, clk Clock) *Register {
|
||||
return &Register{users: users, hasher: hasher, clk: clk}
|
||||
}
|
||||
|
||||
// Execute registers email/password and returns the persisted user. The returned
|
||||
// user carries no plaintext password, only its hash.
|
||||
func (r *Register) Execute(ctx context.Context, email, password string) (*domain.User, error) {
|
||||
if len(password) < minPasswordLen {
|
||||
return nil, ErrPasswordTooShort
|
||||
}
|
||||
if len(password) > maxPasswordLen {
|
||||
return nil, ErrPasswordTooLong
|
||||
}
|
||||
|
||||
hash, err := r.hasher.Hash(password)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// NewUser normalises the email and enforces its shape; it returns a domain
|
||||
// validation error the transport layer maps to 400.
|
||||
u, err := domain.NewUser(email, hash, r.clk.Now())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
if err := r.users.Insert(ctx, u); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
@@ -0,0 +1,133 @@
|
||||
package usecase_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/emil28092005/SciMesh/users/internal/auth"
|
||||
"github.com/emil28092005/SciMesh/users/internal/domain"
|
||||
"github.com/emil28092005/SciMesh/users/internal/memstore"
|
||||
"github.com/emil28092005/SciMesh/users/internal/usecase"
|
||||
)
|
||||
|
||||
const secret = "usecase-test-secret-32-bytes-long!!!"
|
||||
|
||||
func newFixtures() (*usecase.Register, *usecase.Login, *memstore.UserRepo) {
|
||||
users := memstore.NewUserRepo()
|
||||
hasher := auth.NewHasher(4) // low cost keeps tests fast
|
||||
clk := memstore.Clock{T: time.Date(2026, 7, 26, 0, 0, 0, 0, time.UTC)}
|
||||
// The issuer uses the real clock (nil): token expiry is validated against
|
||||
// wall-clock time, so a fixed issue-time would make tokens instantly stale.
|
||||
issuer := auth.NewIssuer(secret, time.Hour, nil)
|
||||
|
||||
reg := usecase.NewRegister(users, hasher, clk)
|
||||
login := usecase.NewLogin(users, hasher, issuer)
|
||||
return reg, login, users
|
||||
}
|
||||
|
||||
func TestRegisterSuccess(t *testing.T) {
|
||||
reg, _, users := newFixtures()
|
||||
|
||||
u, err := reg.Execute(context.Background(), "Alice@Example.com", "password123")
|
||||
if err != nil {
|
||||
t.Fatalf("register: %v", err)
|
||||
}
|
||||
if u.Email != "alice@example.com" {
|
||||
t.Errorf("email not normalised: %q", u.Email)
|
||||
}
|
||||
if u.Role != domain.RoleUser {
|
||||
t.Errorf("role = %q, want user", u.Role)
|
||||
}
|
||||
if strings.Contains(u.PasswordHash, "password123") {
|
||||
t.Error("password stored in cleartext")
|
||||
}
|
||||
if _, err := users.GetByEmail(context.Background(), "alice@example.com"); err != nil {
|
||||
t.Errorf("user not persisted: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterDuplicateEmail(t *testing.T) {
|
||||
reg, _, _ := newFixtures()
|
||||
ctx := context.Background()
|
||||
|
||||
if _, err := reg.Execute(ctx, "dup@example.com", "password123"); err != nil {
|
||||
t.Fatalf("first register: %v", err)
|
||||
}
|
||||
_, err := reg.Execute(ctx, "Dup@example.com", "password123") // different case, same email
|
||||
if !errors.Is(err, usecase.ErrEmailExists) {
|
||||
t.Errorf("got %v, want ErrEmailExists", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterPasswordPolicy(t *testing.T) {
|
||||
reg, _, _ := newFixtures()
|
||||
ctx := context.Background()
|
||||
|
||||
if _, err := reg.Execute(ctx, "a@b.com", "short"); !errors.Is(err, usecase.ErrPasswordTooShort) {
|
||||
t.Errorf("short password: got %v", err)
|
||||
}
|
||||
long := strings.Repeat("x", 73)
|
||||
if _, err := reg.Execute(ctx, "a@b.com", long); !errors.Is(err, usecase.ErrPasswordTooLong) {
|
||||
t.Errorf("long password: got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterInvalidEmail(t *testing.T) {
|
||||
reg, _, _ := newFixtures()
|
||||
_, err := reg.Execute(context.Background(), "not-an-email", "password123")
|
||||
if !errors.Is(err, domain.ErrInvalidEmail) {
|
||||
t.Errorf("got %v, want ErrInvalidEmail", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginSuccess(t *testing.T) {
|
||||
reg, login, _ := newFixtures()
|
||||
ctx := context.Background()
|
||||
if _, err := reg.Execute(ctx, "user@example.com", "password123"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
token, u, err := login.Execute(ctx, "User@Example.com", "password123")
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
if token == "" {
|
||||
t.Error("empty token")
|
||||
}
|
||||
if u.Email != "user@example.com" {
|
||||
t.Errorf("wrong user returned: %q", u.Email)
|
||||
}
|
||||
|
||||
// The token must verify and carry this user's id.
|
||||
claims, err := auth.NewIssuer(secret, time.Hour, nil).Verify(token)
|
||||
if err != nil {
|
||||
t.Fatalf("issued token does not verify: %v", err)
|
||||
}
|
||||
if claims.Subject != u.ID.String() {
|
||||
t.Errorf("token sub = %q, want %q", claims.Subject, u.ID.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginWrongPassword(t *testing.T) {
|
||||
reg, login, _ := newFixtures()
|
||||
ctx := context.Background()
|
||||
if _, err := reg.Execute(ctx, "user@example.com", "password123"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
_, _, err := login.Execute(ctx, "user@example.com", "wrongpass1")
|
||||
if !errors.Is(err, usecase.ErrInvalidCredentials) {
|
||||
t.Errorf("got %v, want ErrInvalidCredentials", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginUnknownEmailIsIndistinguishable(t *testing.T) {
|
||||
_, login, _ := newFixtures()
|
||||
_, _, err := login.Execute(context.Background(), "ghost@example.com", "password123")
|
||||
if !errors.Is(err, usecase.ErrInvalidCredentials) {
|
||||
t.Errorf("unknown email must return ErrInvalidCredentials, got %v", err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user