Implement Sprint 2 security enhancements
CI/CD Pipeline / Test (push) Has been cancelled
CI/CD Pipeline / Lint (push) Has been cancelled
CI/CD Pipeline / Build and Push Docker Image (push) Has been cancelled

Add user-based rate limiting middleware (100 req/min) for authenticated endpoints using httprate library. Implement common password validation blocking 50+ weak passwords. Improve test coverage from 38.3% to 44.6% with comprehensive refresh token tests.

Security improvements:
- Rate limiting by user ID for /habits and /stats endpoints
- X-RateLimit-Limit header in responses
- Common password blacklist in password validation
- Refresh token test suite with 5 scenarios (valid, invalid, expired, revoked, empty)
This commit is contained in:
2025-11-27 10:33:01 +01:00
parent bbe0757ab6
commit 7cb2756b67
17 changed files with 350 additions and 71 deletions
@@ -537,5 +537,3 @@ func TestMarkHabitHandler_CounterFirstMarkWithNegative(t *testing.T) {
t.Fatalf("Expected no error, got %v", err)
}
}
@@ -137,4 +137,3 @@ func TestUnmarkHabitHandler_ReturnsErrorWhenEntryNotFound(t *testing.T) {
t.Errorf("Expected ErrNotFound for missing entry, got %v", err)
}
}
@@ -278,4 +278,3 @@ func TestGetHabitEntriesHandler_RequiresPaginationWithLongDateRange(t *testing.T
t.Errorf("Expected ErrInvalidInput for date range > 1 year without pagination, got %v", err)
}
}
@@ -0,0 +1,176 @@
package queries
import (
"context"
"testing"
"time"
"apocapoc-api/internal/domain/entities"
"apocapoc-api/internal/shared/errors"
)
type mockRefreshTokenRepository struct {
findByTokenFunc func(ctx context.Context, token string) (*entities.RefreshToken, error)
}
func (m *mockRefreshTokenRepository) Create(ctx context.Context, token *entities.RefreshToken) error {
return nil
}
func (m *mockRefreshTokenRepository) FindByToken(ctx context.Context, token string) (*entities.RefreshToken, error) {
if m.findByTokenFunc != nil {
return m.findByTokenFunc(ctx, token)
}
return nil, errors.ErrNotFound
}
func (m *mockRefreshTokenRepository) FindByUserID(ctx context.Context, userID string) ([]*entities.RefreshToken, error) {
return nil, nil
}
func (m *mockRefreshTokenRepository) RevokeByToken(ctx context.Context, token string) error {
return nil
}
func (m *mockRefreshTokenRepository) RevokeAllByUserID(ctx context.Context, userID string) error {
return nil
}
func (m *mockRefreshTokenRepository) DeleteExpired(ctx context.Context) error {
return nil
}
type mockUserRepositoryForRefresh struct {
findByIDFunc func(ctx context.Context, id string) (*entities.User, error)
}
func (m *mockUserRepositoryForRefresh) Create(ctx context.Context, user *entities.User) error {
return nil
}
func (m *mockUserRepositoryForRefresh) FindByID(ctx context.Context, id string) (*entities.User, error) {
if m.findByIDFunc != nil {
return m.findByIDFunc(ctx, id)
}
return nil, errors.ErrNotFound
}
func (m *mockUserRepositoryForRefresh) FindByEmail(ctx context.Context, email string) (*entities.User, error) {
return nil, nil
}
func (m *mockUserRepositoryForRefresh) Update(ctx context.Context, user *entities.User) error {
return nil
}
func TestRefreshTokenHandler_Success(t *testing.T) {
refreshTokenRepo := &mockRefreshTokenRepository{
findByTokenFunc: func(ctx context.Context, token string) (*entities.RefreshToken, error) {
return entities.NewRefreshToken("user-123", token, time.Now().Add(24*time.Hour)), nil
},
}
userRepo := &mockUserRepositoryForRefresh{
findByIDFunc: func(ctx context.Context, id string) (*entities.User, error) {
user := entities.NewUser("test@example.com", "hash", "UTC")
user.ID = id
return user, nil
},
}
handler := NewRefreshTokenHandler(refreshTokenRepo, userRepo)
query := RefreshTokenQuery{
RefreshToken: "valid-token",
}
result, err := handler.Handle(context.Background(), query)
if err != nil {
t.Fatalf("Expected no error, got %v", err)
}
if result.UserID != "user-123" {
t.Errorf("Expected userID 'user-123', got %s", result.UserID)
}
if result.Email != "test@example.com" {
t.Errorf("Expected email 'test@example.com', got %s", result.Email)
}
}
func TestRefreshTokenHandler_InvalidToken(t *testing.T) {
refreshTokenRepo := &mockRefreshTokenRepository{}
userRepo := &mockUserRepositoryForRefresh{}
handler := NewRefreshTokenHandler(refreshTokenRepo, userRepo)
query := RefreshTokenQuery{
RefreshToken: "invalid-token",
}
_, err := handler.Handle(context.Background(), query)
if err != errors.ErrNotFound {
t.Errorf("Expected ErrNotFound, got %v", err)
}
}
func TestRefreshTokenHandler_ExpiredToken(t *testing.T) {
refreshTokenRepo := &mockRefreshTokenRepository{
findByTokenFunc: func(ctx context.Context, token string) (*entities.RefreshToken, error) {
return entities.NewRefreshToken("user-123", token, time.Now().Add(-1*time.Hour)), nil
},
}
userRepo := &mockUserRepositoryForRefresh{}
handler := NewRefreshTokenHandler(refreshTokenRepo, userRepo)
query := RefreshTokenQuery{
RefreshToken: "expired-token",
}
_, err := handler.Handle(context.Background(), query)
if err != errors.ErrNotFound {
t.Errorf("Expected ErrNotFound for expired token, got %v", err)
}
}
func TestRefreshTokenHandler_RevokedToken(t *testing.T) {
refreshToken := entities.NewRefreshToken("user-123", "revoked-token", time.Now().Add(24*time.Hour))
refreshToken.Revoke()
refreshTokenRepo := &mockRefreshTokenRepository{
findByTokenFunc: func(ctx context.Context, token string) (*entities.RefreshToken, error) {
return refreshToken, nil
},
}
userRepo := &mockUserRepositoryForRefresh{}
handler := NewRefreshTokenHandler(refreshTokenRepo, userRepo)
query := RefreshTokenQuery{
RefreshToken: "revoked-token",
}
_, err := handler.Handle(context.Background(), query)
if err != errors.ErrNotFound {
t.Errorf("Expected ErrNotFound for revoked token, got %v", err)
}
}
func TestRefreshTokenHandler_EmptyToken(t *testing.T) {
refreshTokenRepo := &mockRefreshTokenRepository{}
userRepo := &mockUserRepositoryForRefresh{}
handler := NewRefreshTokenHandler(refreshTokenRepo, userRepo)
query := RefreshTokenQuery{
RefreshToken: "",
}
_, err := handler.Handle(context.Background(), query)
if err != errors.ErrInvalidInput {
t.Errorf("Expected ErrInvalidInput, got %v", err)
}
}
@@ -43,4 +43,3 @@ func TestNewHabitEntry_BooleanHabit(t *testing.T) {
t.Error("Value should be nil for boolean habit")
}
}
@@ -0,0 +1,38 @@
package http
import (
"net/http"
"strconv"
"time"
"apocapoc-api/internal/infrastructure/auth"
"github.com/go-chi/httprate"
)
func RateLimitByUser(jwtService *auth.JWTService, requestsPerMinute int, duration time.Duration) func(http.Handler) http.Handler {
limiter := httprate.NewRateLimiter(
requestsPerMinute,
duration,
httprate.WithKeyFuncs(func(r *http.Request) (string, error) {
userID, ok := GetUserIDFromContext(r.Context())
if !ok {
return r.RemoteAddr, nil
}
return "user:" + userID, nil
}),
httprate.WithLimitHandler(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusTooManyRequests)
w.Write([]byte(`{"error":"Rate limit exceeded. Please try again later."}`))
}),
)
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-RateLimit-Limit", strconv.Itoa(requestsPerMinute))
limiter.Handler(next).ServeHTTP(w, r)
})
}
}
+3
View File
@@ -46,6 +46,8 @@ func NewRouter(corsOrigins string, habitHandlers *HabitHandlers, authHandlers *A
r.Route("/api/v1/habits", func(r chi.Router) {
r.Use(AuthMiddleware(jwtService))
r.Use(RateLimitByUser(jwtService, 100, 1*time.Minute))
r.Post("/", habitHandlers.CreateHabit)
r.Get("/", habitHandlers.GetUserHabits)
r.Get("/today", habitHandlers.GetTodaysHabits)
@@ -59,6 +61,7 @@ func NewRouter(corsOrigins string, habitHandlers *HabitHandlers, authHandlers *A
r.Route("/api/v1/stats", func(r chi.Router) {
r.Use(AuthMiddleware(jwtService))
r.Use(RateLimitByUser(jwtService, 100, 1*time.Minute))
r.Get("/habits/{id}", statsHandlers.GetHabitStats)
})
@@ -0,0 +1,63 @@
package validation
import "strings"
var commonPasswords = map[string]bool{
"123456": true,
"password": true,
"123456789": true,
"12345678": true,
"12345": true,
"1234567": true,
"password1": true,
"123123": true,
"1234567890": true,
"000000": true,
"abc123": true,
"1234": true,
"qwerty": true,
"111111": true,
"123321": true,
"dragon": true,
"master": true,
"monkey": true,
"letmein": true,
"login": true,
"princess": true,
"qwertyuiop": true,
"solo": true,
"passw0rd": true,
"starwars": true,
"iloveyou": true,
"welcome": true,
"admin": true,
"sunshine": true,
"password123": true,
"123qwe": true,
"654321": true,
"superman": true,
"1qaz2wsx": true,
"trustno1": true,
"charlie": true,
"666666": true,
"qazwsx": true,
"freedom": true,
"football": true,
"baseball": true,
"whatever": true,
"jordan": true,
"killer": true,
"summer": true,
"hockey": true,
"bailey": true,
"shadow": true,
"master123": true,
"ninja": true,
"mustang": true,
"password!": true,
}
func IsCommonPassword(password string) bool {
lower := strings.ToLower(password)
return commonPasswords[lower]
}
+4
View File
@@ -92,6 +92,10 @@ func ValidatePassword(password string) error {
return ValidationError{Field: "password", Message: "password must contain at least one special character"}
}
if IsCommonPassword(password) {
return ValidationError{Field: "password", Message: "password is too common, please choose a more secure password"}
}
return nil
}