Add JWT authentication #7

Merged
robin merged 3 commits from dev into main 2026-07-21 13:49:49 +02:00
12 changed files with 296 additions and 10 deletions
Showing only changes of commit d9bfcea420 - Show all commits
+1
View File
@@ -1,2 +1,3 @@
DTTMR_OTLP_ENDPOINT=172.17.0.1:4317 DTTMR_OTLP_ENDPOINT=172.17.0.1:4317
DTTMR_DATABASE_URL=postgres://postgres:postgres@172.17.0.1:5432/postgres?sslmode=disable&timezone=utc DTTMR_DATABASE_URL=postgres://postgres:postgres@172.17.0.1:5432/postgres?sslmode=disable&timezone=utc
DTTMR_JWT_SECRET=V*$#Jt9OlYW0gEB6PyUU$qLYbJ^NC7LZ
+5 -4
View File
@@ -76,7 +76,7 @@ func run(serviceName string, serviceVersion string) error {
slog.Error("failed to run migrations", slog.Any("error", err)) slog.Error("failed to run migrations", slog.Any("error", err))
} }
srv := makeServer(db, cfg.Port) srv := makeServer(db, cfg)
go func() { go func() {
slog.Info("starting http server", "addr", srv.Addr) slog.Info("starting http server", "addr", srv.Addr)
if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) {
@@ -112,14 +112,15 @@ func setupLogging() {
slog.SetDefault(logger) slog.SetDefault(logger)
} }
func makeServer(db *sql.DB, port int) *http.Server { func makeServer(db *sql.DB, cfg *config.Config) *http.Server {
routerConfig := router.Config{ routerConfig := router.Config{
Database: db, Database: db,
JWTSecret: cfg.JWTSecret,
} }
mux := router.NewMux(routerConfig) mux := router.NewMux(routerConfig)
srv := &http.Server{ srv := &http.Server{
Addr: fmt.Sprintf(":%d", port), Addr: fmt.Sprintf(":%d", cfg.Port),
Handler: mux, Handler: mux,
ReadTimeout: 5 * time.Second, ReadTimeout: 5 * time.Second,
WriteTimeout: 10 * time.Second, WriteTimeout: 10 * time.Second,
+1
View File
@@ -22,6 +22,7 @@ require (
github.com/felixge/httpsnoop v1.1.0 // indirect github.com/felixge/httpsnoop v1.1.0 // indirect
github.com/go-logr/logr v1.4.4 // indirect github.com/go-logr/logr v1.4.4 // indirect
github.com/go-logr/stdr v1.2.2 // indirect github.com/go-logr/stdr v1.2.2 // indirect
github.com/golang-jwt/jwt/v5 v5.3.1 // indirect
github.com/google/uuid v1.6.0 // indirect github.com/google/uuid v1.6.0 // indirect
github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect
github.com/jackc/pgpassfile v1.0.0 // indirect github.com/jackc/pgpassfile v1.0.0 // indirect
+38
View File
@@ -0,0 +1,38 @@
package handler
import (
"log/slog"
"net/http"
"github.com/robindittmar/dttmr-api/internal/api/request"
"github.com/robindittmar/dttmr-api/internal/api/response"
"github.com/robindittmar/dttmr-api/internal/domain"
)
type AuthHandler struct {
AuthService *domain.AuthService
}
func NewAuthHandler(authService *domain.AuthService) *AuthHandler {
return &AuthHandler{AuthService: authService}
}
func (h *AuthHandler) Login(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
payload, err := request.DecodeLogin(r)
if err != nil {
slog.ErrorContext(ctx, "failed to decode login payload", slog.Any("error", err))
response.Error(ctx, w, http.StatusBadRequest, "failed to decode request body")
return
}
token, err := h.AuthService.Login(ctx, payload.Email, payload.Password)
if err != nil {
slog.ErrorContext(ctx, "failed to login", slog.Any("error", err))
response.Error(ctx, w, http.StatusInternalServerError, "failed to login")
return
}
response.JSON(ctx, w, http.StatusOK, token)
}
+42
View File
@@ -0,0 +1,42 @@
package middleware
import (
"context"
"log/slog"
"net/http"
"strings"
"github.com/robindittmar/dttmr-api/internal/api/response"
"github.com/robindittmar/dttmr-api/internal/domain"
)
func WithJWT(authService *domain.AuthService) func(http.HandlerFunc) http.HandlerFunc {
return func(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
ctx := r.Context()
authHeader := r.Header.Get("Authorization")
if authHeader == "" {
slog.ErrorContext(ctx, "missing authorization header")
response.Error(ctx, w, http.StatusUnauthorized, "missing authorization header")
return
}
parts := strings.SplitN(authHeader, " ", 2)
if len(parts) != 2 || !strings.EqualFold(parts[0], "Bearer") {
slog.ErrorContext(ctx, "invalid authorization header")
response.Error(ctx, w, http.StatusUnauthorized, "invalid authorization header")
return
}
authContext, err := authService.ParseToken(ctx, parts[1])
if err != nil {
slog.ErrorContext(ctx, "invalid token", slog.Any("error", err))
response.Error(ctx, w, http.StatusUnauthorized, "invalid or expired token")
return
}
ctx = context.WithValue(ctx, domain.AuthContextKey, authContext)
next.ServeHTTP(w, r.WithContext(ctx))
}
}
}
+25
View File
@@ -0,0 +1,25 @@
package request
import (
"encoding/json"
"fmt"
"net/http"
)
type LoginPayload struct {
Email string `json:"email"`
Password string `json:"password"`
}
func DecodeLogin(r *http.Request) (LoginPayload, error) {
var payload LoginPayload
decoder := json.NewDecoder(r.Body)
decoder.DisallowUnknownFields()
if err := decoder.Decode(&payload); err != nil {
return payload, fmt.Errorf("error decoding login payload: %w", err)
}
return payload, nil
}
+11 -4
View File
@@ -11,7 +11,8 @@ import (
) )
type Config struct { type Config struct {
Database *sql.DB Database *sql.DB
JWTSecret string
} }
func NewMux(cfg Config) http.Handler { func NewMux(cfg Config) http.Handler {
@@ -19,18 +20,24 @@ func NewMux(cfg Config) http.Handler {
userService := domain.NewUserService(userRepo) userService := domain.NewUserService(userRepo)
userHandler := handler.NewUserHandler(userService) userHandler := handler.NewUserHandler(userService)
authRepo := repository.NewAuthRepo(cfg.Database)
authService := domain.NewAuthService(authRepo, []byte(cfg.JWTSecret))
authHandler := handler.NewAuthHandler(authService)
listRepo := repository.NewListRepo(cfg.Database) listRepo := repository.NewListRepo(cfg.Database)
listService := domain.NewListService(listRepo) listService := domain.NewListService(listRepo)
listHandler := handler.NewListHandler(listService) listHandler := handler.NewListHandler(listService)
mux := http.NewServeMux() protected := middleware.WithJWT(authService)
mux := http.NewServeMux()
mux.HandleFunc("/", handler.DefaultHandler) mux.HandleFunc("/", handler.DefaultHandler)
mux.HandleFunc("GET /health", handler.HealthHandler) mux.HandleFunc("GET /health", handler.HealthHandler)
mux.HandleFunc("POST /login", authHandler.Login)
mux.HandleFunc("POST /users", userHandler.CreateUser) mux.Handle("POST /users", protected(userHandler.CreateUser))
mux.HandleFunc("POST /lists", listHandler.CreateList) mux.Handle("POST /lists", protected(listHandler.CreateList))
var httpHandler http.Handler = mux var httpHandler http.Handler = mux
httpHandler = middleware.WithMaxBytes(1024 * 64)(httpHandler) httpHandler = middleware.WithMaxBytes(1024 * 64)(httpHandler)
+4
View File
@@ -12,6 +12,7 @@ type Config struct {
Port int Port int
OTLPEndpoint string OTLPEndpoint string
DatabaseURL string DatabaseURL string
JWTSecret string
} }
func Load() *Config { func Load() *Config {
@@ -19,6 +20,7 @@ func Load() *Config {
portFlag := flag.Int("port", 8080, "port to listen on") portFlag := flag.Int("port", 8080, "port to listen on")
otlpEndpointFlag := flag.String("otlp-endpoint", "localhost:4317", "otlp endpoint") otlpEndpointFlag := flag.String("otlp-endpoint", "localhost:4317", "otlp endpoint")
databaseUrlFlag := flag.String("database-url", "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable&timezone=utc", "database connection string") databaseUrlFlag := flag.String("database-url", "postgres://postgres:postgres@localhost:5432/postgres?sslmode=disable&timezone=utc", "database connection string")
jwtSecretFlag := flag.String("jwt-secret", "5!zM8k@wC0Y5jgrbS8xLC0gW9k7dLaeI", "JWT secret")
flag.Parse() flag.Parse()
@@ -27,12 +29,14 @@ func Load() *Config {
Port: *portFlag, Port: *portFlag,
OTLPEndpoint: *otlpEndpointFlag, OTLPEndpoint: *otlpEndpointFlag,
DatabaseURL: *databaseUrlFlag, DatabaseURL: *databaseUrlFlag,
JWTSecret: *jwtSecretFlag,
} }
assignStringFromEnv("DTTMR_ENVIRONMENT", &cfg.Environment) assignStringFromEnv("DTTMR_ENVIRONMENT", &cfg.Environment)
assignIntFromEnv("DTTMR_PORT", &cfg.Port) assignIntFromEnv("DTTMR_PORT", &cfg.Port)
assignStringFromEnv("DTTMR_OTLP_ENDPOINT", &cfg.OTLPEndpoint) assignStringFromEnv("DTTMR_OTLP_ENDPOINT", &cfg.OTLPEndpoint)
assignStringFromEnv("DTTMR_DATABASE_URL", &cfg.DatabaseURL) assignStringFromEnv("DTTMR_DATABASE_URL", &cfg.DatabaseURL)
assignStringFromEnv("DTTMR_JWT_SECRET", &cfg.JWTSecret)
return cfg return cfg
} }
+136
View File
@@ -0,0 +1,136 @@
package domain
import (
"context"
"errors"
"fmt"
"log/slog"
"time"
"github.com/golang-jwt/jwt/v5"
"golang.org/x/crypto/bcrypt"
)
type AuthRepository interface {
GetByEmail(ctx context.Context, email string) (*AuthUser, error)
}
type AuthService struct {
repo AuthRepository
jwtSecret []byte
}
type AuthToken struct {
Token string `json:"token"`
}
type AuthUser struct {
ID string
Email string
Name string
PasswordHash string
}
type AuthContext struct {
UserID string
Email string
Name string
}
func GetAuthContext(ctx context.Context) (*AuthContext, error) {
v := ctx.Value(AuthContextKey)
if v == nil {
return nil, errors.New("no auth context")
}
ac, ok := v.(*AuthContext)
if !ok {
return nil, errors.New("invalid auth context")
}
return ac, nil
}
func NewAuthService(r AuthRepository, jwtSecret []byte) *AuthService {
return &AuthService{repo: r, jwtSecret: jwtSecret}
}
func (s *AuthService) authenticate(ctx context.Context, email string, password string) (*AuthUser, error) {
user, err := s.repo.GetByEmail(ctx, email)
if err != nil {
return user, err
}
err = bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password))
if err != nil {
if errors.Is(err, bcrypt.ErrMismatchedHashAndPassword) {
return user, errors.New("invalid email or password")
}
return user, err
}
return user, nil
}
func (s *AuthService) Login(ctx context.Context, email string, password string) (AuthToken, error) {
var authToken AuthToken
user, err := s.authenticate(ctx, email, password)
if err != nil {
return authToken, err
}
authToken.Token, err = s.GenerateToken(user)
if err != nil {
return authToken, err
}
return authToken, nil
}
type JWTClaims struct {
UserID string `json:"user_id"`
Email string `json:"email"`
Name string `json:"name"`
jwt.RegisteredClaims
}
type contextKey string
const AuthContextKey = contextKey("auth")
func (s *AuthService) GenerateToken(authUser *AuthUser) (string, error) {
claims := JWTClaims{
UserID: authUser.ID,
Email: authUser.Email,
Name: authUser.Name,
}
claims.ExpiresAt = jwt.NewNumericDate(time.Now().Add(time.Hour * 24))
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
tokenString, err := token.SignedString(s.jwtSecret)
if err != nil {
return "", err
}
return tokenString, nil
}
func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (*AuthContext, error) {
claims := &JWTClaims{}
token, err := jwt.ParseWithClaims(tokenString, claims, func(token *jwt.Token) (any, error) {
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"])
}
return s.jwtSecret, nil
})
if err != nil || !token.Valid {
slog.ErrorContext(ctx, "invalid or expired token", slog.Any("token", token))
return nil, err
}
return &AuthContext{
UserID: claims.UserID,
Email: claims.Email,
Name: claims.Name,
}, nil
}
+31
View File
@@ -0,0 +1,31 @@
package repository
import (
"context"
"database/sql"
"fmt"
"github.com/robindittmar/dttmr-api/internal/domain"
)
type AuthRepo struct {
db *sql.DB
}
func NewAuthRepo(db *sql.DB) *AuthRepo {
return &AuthRepo{db: db}
}
func (r *AuthRepo) GetByEmail(ctx context.Context, email string) (*domain.AuthUser, error) {
user := &domain.AuthUser{}
err := r.db.QueryRowContext(ctx,
"SELECT id, email, name, password_hash FROM users WHERE email = $1",
email,
).Scan(&user.ID, &user.Email, &user.Name, &user.PasswordHash)
if err != nil {
return nil, fmt.Errorf("failed to get user: %w", err)
}
return user, nil
}
+1 -1
View File
@@ -55,5 +55,5 @@ func (r *ListRepo) CreateList(ctx context.Context, name string, userIDs []string
return nil, fmt.Errorf("commit transaction: %w", err) return nil, fmt.Errorf("commit transaction: %w", err)
} }
return list, err return list, nil
} }
+1 -1
View File
@@ -37,5 +37,5 @@ func (r *UserRepo) CreateUser(ctx context.Context, email string, name string, pa
return nil, fmt.Errorf("commit transaction: %w", err) return nil, fmt.Errorf("commit transaction: %w", err)
} }
return user, err return user, nil
} }