Files
bookhoard/cmd/server/main.go
T
john-okeefe a3aa9f67ac feat: add Kobo device sync support and fix device route protection
- Add Kobo sync handler with markup, bookmark, analytics, and initialization endpoints
- Add Kobo integration tests and Bruno API test collection
- Move device approve/reject routes from public to protected routes
- Enhance test infrastructure with DATABASE_URL support and helper functions
- Fix device GetDevice handler nil pointer handling
- Clean up test reports and session files
2026-01-30 23:58:34 -05:00

384 lines
12 KiB
Go

package main
import (
"bookmann/internal/config"
"bookmann/internal/database"
"bookmann/internal/handlers"
"bookmann/internal/middleware"
ratelimit "bookmann/internal/middleware"
"bookmann/internal/sync"
"bookmann/templates"
"bytes"
"context"
"log"
"net/http"
"strings"
"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)
deviceHandler := handlers.NewDeviceHandler(queries, cfg.JWTSecret, cfg)
deviceAuthMiddleware := middleware.NewDeviceAuthMiddleware(queries)
// Create WebSocket connection manager
connManager := sync.NewConnectionManager()
connManager.StartCleanupTask()
koreaderHandler := handlers.NewKOReaderHandler(queries, connManager)
wsHandler := handlers.NewWSHandler(queries, connManager, cfg.JWTSecret, deviceAuthMiddleware)
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)
// Setup ebook handler routes first (so we can use it for library scan)
h := handlers.SetupRoutes(protected, queries, connManager)
// Public library types endpoint (no authentication required)
e.GET("/api/libraries/types", libraryHandler.GetLibraryTypes)
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")
// 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)
adminLibrary.POST("/:id/scan", func(c echo.Context) error {
libraryID := c.Param("id")
scanReq := map[string]interface{}{
"library_id": libraryID,
}
c.Set("scan_request", scanReq)
return h.ScanEbooks(c)
})
adminLibrary.GET("/:id/media-items", func(c echo.Context) error {
libraryID := c.Param("id")
c.QueryParams().Set("library_id", libraryID)
return h.ListMediaItems(c)
})
// 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
// Device management routes (public - for registration)
e.POST("/api/devices/register", deviceHandler.InitiateRegistration)
e.POST("/api/devices/register/status", deviceHandler.CheckRegistrationStatus)
// KOReader sync routes (device authentication required)
koreaderSync := e.Group("/api/sync/koreader")
koreaderSync.POST("/progress", deviceAuthMiddleware.Authenticate(koreaderHandler.SyncProgress))
koreaderSync.GET("/metadata/:uuid", deviceAuthMiddleware.Authenticate(koreaderHandler.GetMetadata))
koreaderSync.GET("/library", deviceAuthMiddleware.Authenticate(koreaderHandler.GetLibrary))
koreaderSync.POST("/bookmarks", deviceAuthMiddleware.Authenticate(koreaderHandler.SyncBookmarks))
// Kobo sync routes (device authentication required)
koboHandler := handlers.NewKoboHandler(queries, connManager)
koboSync := e.Group("/api/sync/kobo")
koboSync.POST("/markup", deviceAuthMiddleware.Authenticate(koboHandler.Markup))
koboSync.POST("/bookmark", deviceAuthMiddleware.Authenticate(koboHandler.Bookmark))
koboSync.POST("/v1/analytics/gettests", deviceAuthMiddleware.Authenticate(koboHandler.AnalyticsGettests))
koboSync.GET("/v1/initialization", deviceAuthMiddleware.Authenticate(koboHandler.Initialization))
// Device management routes (protected - require user auth)
devices := protected.Group("/devices")
devices.GET("", deviceHandler.ListDevices)
devices.GET("/:id", deviceHandler.GetDevice)
devices.PUT("/:id", deviceHandler.UpdateDevice)
devices.DELETE("/:id", deviceHandler.DeleteDevice)
devices.GET("/pending", deviceHandler.ListPendingRegistrations)
devices.GET("/approve/:registration_id", deviceHandler.ApproveDevice)
devices.POST("/reject/:registration_id", deviceHandler.RejectDevice)
// WebSocket endpoint for real-time sync
e.GET("/ws/sync", wsHandler.HandleWebSocket)
// Static files
e.Static("/static", "web/static")
// 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)
}
}()
// Bookshelf route (protected) - new default for logged-in users
protected.GET("/bookshelf", 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.BookShelf(user).Render(c.Request().Context(), &buf)
if err != nil {
return err
}
return c.HTML(http.StatusOK, buf.String())
})
// Direct /bookshelf route (protected)
e.GET("/bookshelf", func(c echo.Context) error {
tokenString := c.Request().Header.Get("Authorization")
if tokenString != "" && strings.HasPrefix(tokenString, "Bearer ") {
tokenString = tokenString[7:]
} else {
// Check for token in cookie
cookie, err := c.Cookie("token")
if err != nil {
return c.Redirect(http.StatusFound, "/login")
}
tokenString = cookie.Value
}
token, err := jwt.Parse(tokenString, func(token *jwt.Token) (interface{}, error) {
return []byte(cfg.JWTSecret), nil
})
if err != nil || !token.Valid {
return c.Redirect(http.StatusFound, "/login")
}
claims := token.Claims.(jwt.MapClaims)
user := templates.User{
ID: claims["user_id"].(string),
Email: claims["user_email"].(string),
Username: claims["user_username"].(string),
Role: claims["user_role"].(string),
}
var buf bytes.Buffer
err = templates.BookShelf(user).Render(c.Request().Context(), &buf)
if err != nil {
return err
}
return c.HTML(http.StatusOK, buf.String())
})
// Dashboard route (protected) - keep for backward compatibility
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))
}