- Add TestMode, RateLimitEnabled, RequestsPerMinute to Config - Add getEnvBool() and getEnvInt() helper functions - Update rate limiter to support enabled/disabled state - Pass test environment variables through docker-compose - Configure rate limiter dynamically in main.go This allows disabling rate limiting for integration testing while maintaining security in production environments.
264 lines
7.9 KiB
Go
264 lines
7.9 KiB
Go
package main
|
|
|
|
import (
|
|
"bookmann/internal/config"
|
|
"bookmann/internal/database"
|
|
"bookmann/internal/handlers"
|
|
ratelimit "bookmann/internal/middleware"
|
|
"bookmann/templates"
|
|
"bytes"
|
|
"context"
|
|
"log"
|
|
"net/http"
|
|
"time"
|
|
|
|
"github.com/go-playground/validator/v10"
|
|
"github.com/golang-jwt/jwt/v5"
|
|
"github.com/google/uuid"
|
|
"github.com/jackc/pgx/v5/pgtype"
|
|
"github.com/jackc/pgx/v5/pgxpool"
|
|
"github.com/labstack/echo-jwt/v4"
|
|
"github.com/labstack/echo/v4"
|
|
echomiddleware "github.com/labstack/echo/v4/middleware"
|
|
)
|
|
|
|
// CustomValidator wraps the go-playground validator
|
|
type CustomValidator struct {
|
|
validator *validator.Validate
|
|
}
|
|
|
|
func (cv *CustomValidator) Validate(i interface{}) error {
|
|
return cv.validator.Struct(i)
|
|
}
|
|
|
|
func main() {
|
|
cfg := config.LoadConfig()
|
|
|
|
dbPool, err := pgxpool.New(context.Background(), cfg.DatabaseURL())
|
|
if err != nil {
|
|
log.Fatal("Failed to connect to database:", err)
|
|
}
|
|
defer dbPool.Close()
|
|
|
|
queries := database.New(dbPool)
|
|
|
|
// Create login attempt tracker: 5 failed attempts = 15 minute lockout
|
|
loginAttemptTracker := ratelimit.NewLoginAttemptTracker(5, 15*time.Minute, 5*time.Minute)
|
|
|
|
authHandler := handlers.NewAuthHandler(queries, cfg.JWTSecret, loginAttemptTracker)
|
|
libraryHandler := handlers.NewLibraryHandler(queries)
|
|
|
|
e := echo.New()
|
|
|
|
// Set up validator
|
|
v := validator.New()
|
|
|
|
// Register custom password complexity validator
|
|
if err := ratelimit.RegisterPasswordValidation(v); err != nil {
|
|
log.Fatal("Failed to register password validator:", err)
|
|
}
|
|
|
|
e.Validator = &CustomValidator{validator: v}
|
|
|
|
// Middleware
|
|
e.Use(echomiddleware.Logger())
|
|
e.Use(echomiddleware.Recover())
|
|
e.Use(echomiddleware.CORS())
|
|
e.Use(ratelimit.RequestTracingMiddleware(cfg))
|
|
|
|
// Rate limiter for auth endpoints
|
|
rateLimiterConfig := ratelimit.RateLimiterConfig{
|
|
Enabled: cfg.RateLimitEnabled,
|
|
RequestsPerMinute: cfg.RequestsPerMinute,
|
|
CleanupInterval: 5 * time.Minute,
|
|
}
|
|
rateLimiter := ratelimit.NewRateLimiter(rateLimiterConfig)
|
|
rateLimitMiddleware := ratelimit.RateLimiterMiddleware(rateLimiter)
|
|
|
|
// Auth routes (no auth required, but rate limited)
|
|
e.POST("/api/auth/register", rateLimitMiddleware(authHandler.Register))
|
|
e.POST("/api/auth/login", rateLimitMiddleware(authHandler.Login))
|
|
|
|
// JWT middleware for protected routes
|
|
jwtMiddleware := echojwt.WithConfig(echojwt.Config{
|
|
SigningKey: []byte(cfg.JWTSecret),
|
|
ContextKey: "user",
|
|
SuccessHandler: func(c echo.Context) {
|
|
token := c.Get("user").(*jwt.Token)
|
|
claims := token.Claims.(jwt.MapClaims)
|
|
c.Set("user_id", claims["user_id"])
|
|
c.Set("user_role", claims["user_role"])
|
|
c.Set("user_email", claims["user_email"])
|
|
c.Set("user_username", claims["user_username"])
|
|
// Parse UUID from string claims
|
|
userIDStr, _ := claims["user_id"].(string)
|
|
userUUID, err := uuid.Parse(userIDStr)
|
|
if err != nil {
|
|
c.JSON(http.StatusBadRequest, map[string]string{"error": "invalid user ID in token"})
|
|
return
|
|
}
|
|
|
|
c.Set("user", database.Users{
|
|
ID: pgtype.UUID{Bytes: [16]byte(userUUID), Valid: true},
|
|
Email: claims["user_email"].(string),
|
|
Username: claims["user_username"].(string),
|
|
Role: claims["user_role"].(string),
|
|
})
|
|
},
|
|
})
|
|
|
|
// Protected routes
|
|
protected := e.Group("/api", jwtMiddleware)
|
|
protected.GET("/auth/profile", authHandler.GetProfile)
|
|
protected.PUT("/auth/profile", authHandler.UpdateProfile)
|
|
protected.POST("/auth/refresh", authHandler.RefreshAccessToken)
|
|
protected.POST("/auth/logout", authHandler.Logout)
|
|
|
|
// Admin-only routes for user and folder management
|
|
admin := protected.Group("/auth", handlers.AdminMiddleware)
|
|
admin.GET("/users", authHandler.ListUsers)
|
|
|
|
// Library management routes
|
|
library := protected.Group("/libraries")
|
|
library.GET("/types", libraryHandler.GetLibraryTypes)
|
|
|
|
// Admin-only library routes
|
|
adminLibrary := library.Group("", handlers.AdminMiddleware)
|
|
adminLibrary.POST("", libraryHandler.CreateLibrary)
|
|
adminLibrary.GET("", libraryHandler.ListLibraries)
|
|
adminLibrary.GET("/:id", libraryHandler.GetLibrary)
|
|
adminLibrary.PUT("/:id", libraryHandler.UpdateLibrary)
|
|
adminLibrary.DELETE("/:id", libraryHandler.DeleteLibrary)
|
|
adminLibrary.POST("/:id/folders", libraryHandler.AddLibraryFolder)
|
|
adminLibrary.GET("/:id/folders", libraryHandler.GetLibraryFolders)
|
|
adminLibrary.DELETE("/:id/folders", libraryHandler.DeleteLibraryFolder)
|
|
adminLibrary.GET("/:id/stats", libraryHandler.GetLibraryStats)
|
|
|
|
// User library visibility control
|
|
protected.POST("/libraries/visibility", libraryHandler.SetLibraryVisibility)
|
|
protected.GET("/libraries/visible", libraryHandler.GetUserVisibleLibraries)
|
|
|
|
protected.DELETE("/auth/account", authHandler.DeleteAccount)
|
|
protected.PUT("/library/scan-settings", authHandler.UpdateScanSettings)
|
|
protected.GET("/library/scan-settings", authHandler.GetScanSettings)
|
|
|
|
// Auth update routes
|
|
authGroup := e.Group("/api/auth", jwtMiddleware)
|
|
authGroup.PUT("/email", authHandler.UpdateEmail)
|
|
authGroup.PUT("/username", authHandler.UpdateUsername)
|
|
authGroup.PUT("/password", authHandler.UpdatePassword)
|
|
authGroup.PUT("/theme", authHandler.UpdateTheme)
|
|
// force rebuild
|
|
|
|
// Static files
|
|
e.Static("/static", "static")
|
|
|
|
// Routes
|
|
h := handlers.SetupRoutes(protected, queries)
|
|
|
|
// Start scheduler for auto-scanning
|
|
go h.StartScheduler()
|
|
defer h.StopScheduler()
|
|
|
|
// Start watch mode for all libraries (background)
|
|
go func() {
|
|
time.Sleep(2 * time.Second) // Wait a bit for server to be ready
|
|
if err := h.StartWatchModeForAllLibraries(context.Background()); err != nil {
|
|
log.Printf("Warning: failed to start watch mode for libraries: %v", err)
|
|
}
|
|
}()
|
|
|
|
// Dashboard route (protected)
|
|
protected.GET("/dashboard", func(c echo.Context) error {
|
|
userID := c.Get("user_id").(string)
|
|
userEmail := c.Get("user_email").(string)
|
|
userUsername := c.Get("user_username").(string)
|
|
userRole := c.Get("user_role").(string)
|
|
|
|
user := templates.User{
|
|
ID: userID,
|
|
Email: userEmail,
|
|
Username: userUsername,
|
|
Role: userRole,
|
|
}
|
|
|
|
var buf bytes.Buffer
|
|
err := templates.Dashboard(user).Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
dummyUser := templates.User{ID: "", Username: "Admin", Email: "admin@example.com"}
|
|
|
|
// Routes
|
|
e.GET("/", func(c echo.Context) error {
|
|
loggedIn := false
|
|
var buf bytes.Buffer
|
|
err := templates.Index(loggedIn).Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
e.GET("/login", func(c echo.Context) error {
|
|
var buf bytes.Buffer
|
|
err := templates.Login().Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
e.GET("/register", func(c echo.Context) error {
|
|
var buf bytes.Buffer
|
|
err := templates.Register().Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
e.GET("/admin", func(c echo.Context) error {
|
|
var buf bytes.Buffer
|
|
err := templates.Admin(dummyUser).Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
e.GET("/admin/", func(c echo.Context) error {
|
|
var buf bytes.Buffer
|
|
err := templates.Admin(dummyUser).Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
e.GET("/admin/profile", func(c echo.Context) error {
|
|
var buf bytes.Buffer
|
|
err := templates.AdminProfile(dummyUser).Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
e.GET("/admin/library", func(c echo.Context) error {
|
|
var buf bytes.Buffer
|
|
err := templates.AdminLibrary(dummyUser).Render(c.Request().Context(), &buf)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
return c.HTML(http.StatusOK, buf.String())
|
|
})
|
|
|
|
// Start server
|
|
log.Printf("Starting server on port %s", cfg.ServerPort)
|
|
e.Logger.Fatal(e.Start(":" + cfg.ServerPort))
|
|
}
|