add users logic

This commit is contained in:
Efremenko Arhip
2026-07-26 14:38:24 +03:00
parent 16db1e41f7
commit a3db1a1e67
65 changed files with 3626 additions and 19 deletions
+67
View File
@@ -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
}
+74
View File
@@ -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")
}
}
+40
View File
@@ -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))
}
+30
View File
@@ -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)")
}
}
+11
View File
@@ -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")
)
+80
View File
@@ -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
}
+62
View File
@@ -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")
}
}
+13
View File
@@ -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() }
+151
View File
@@ -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
}
+65
View File
@@ -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)
},
)
}
+65
View File
@@ -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 }
+74
View File
@@ -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)
}
}
}
}
+66
View File
@@ -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)
}
}
+82
View File
@@ -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)
}
}
+83
View File
@@ -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
}
+41
View File
@@ -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),
}
}
+61
View File
@@ -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"
}
}
+85
View File
@@ -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
}
+133
View File
@@ -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
}
+41
View File
@@ -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)
}
}
+18
View File
@@ -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")
)
+42
View File
@@ -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
}
+41
View File
@@ -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
}
+56
View File
@@ -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
}
+133
View File
@@ -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)
}
}