package middleware import ( "bookhoard/internal/database" "context" "net/http" "strconv" "strings" "github.com/google/uuid" "github.com/jackc/pgx/v5/pgtype" "github.com/labstack/echo/v4" ) type DeviceContext struct { ID uuid.UUID UserID uuid.UUID DeviceName string DeviceType string DeviceIdentifier string SyncEnabled bool AutoSync bool } type DeviceAuthMiddleware struct { db *database.Queries rateLimiter *DeviceRateLimiter } func NewDeviceAuthMiddleware(db *database.Queries) *DeviceAuthMiddleware { return &DeviceAuthMiddleware{ db: db, rateLimiter: NewDeviceRateLimiter(), } } func (m *DeviceAuthMiddleware) Authenticate(next echo.HandlerFunc) echo.HandlerFunc { return func(c echo.Context) error { authHeader := c.Request().Header.Get("Authorization") if authHeader == "" { return c.JSON(http.StatusUnauthorized, map[string]string{ "error": "missing authorization header", }) } if !strings.HasPrefix(authHeader, "Bearer ") { return c.JSON(http.StatusUnauthorized, map[string]string{ "error": "invalid authorization header format", }) } token := strings.TrimPrefix(authHeader, "Bearer ") device, err := m.db.GetDeviceByAuthToken(c.Request().Context(), token) if err != nil { return c.JSON(http.StatusUnauthorized, map[string]string{ "error": "invalid device token", }) } if !device.SyncEnabled.Bool || !device.SyncEnabled.Valid { return c.JSON(http.StatusForbidden, map[string]string{ "error": "device sync is disabled", }) } requestType := m.getRequestType(c.Request().URL.Path) deviceUUID := uuid.UUID(device.ID.Bytes) deviceID := deviceUUID.String() config := DeviceRateLimitConfig{ SyncRequestsPerMinute: 60, ProgressUpdatesPerMinute: 120, MetadataRequestsPerMinute: 30, } if !m.rateLimiter.CheckRateLimit(deviceID, requestType, config) { remaining := m.rateLimiter.GetRemainingRequests(deviceID, requestType, config) c.Response().Header().Set("X-RateLimit-Limit", "60") c.Response().Header().Set("X-RateLimit-Remaining", strconv.Itoa(remaining)) c.Response().Header().Set("X-RateLimit-Reset", "60") return c.JSON(http.StatusTooManyRequests, map[string]string{ "error": "rate limit exceeded", "message": "Too many requests", "remaining": strconv.Itoa(remaining), }) } remaining := m.rateLimiter.GetRemainingRequests(deviceID, requestType, config) c.Response().Header().Set("X-RateLimit-Limit", "60") c.Response().Header().Set("X-RateLimit-Remaining", strconv.Itoa(remaining)) ctx := DeviceContext{ ID: device.ID.Bytes, UserID: device.UserID.Bytes, DeviceName: device.DeviceName, DeviceType: device.DeviceType, DeviceIdentifier: device.DeviceIdentifier, SyncEnabled: device.SyncEnabled.Bool && device.SyncEnabled.Valid, AutoSync: device.AutoSync.Bool && device.AutoSync.Valid, } c.Set("device", device) c.Set("device_ctx", ctx) c.Set("device_id", device.ID.Bytes) return next(c) } } func (m *DeviceAuthMiddleware) getRequestType(path string) string { if strings.Contains(path, "/progress") { return "progress" } if strings.Contains(path, "/metadata") || strings.Contains(path, "/library") { return "metadata" } return "sync" } func (m *DeviceAuthMiddleware) RequirePermission(permission string) echo.MiddlewareFunc { return func(next echo.HandlerFunc) echo.HandlerFunc { return func(c echo.Context) error { device, ok := c.Get("device").(database.Devices) if !ok { return c.JSON(http.StatusUnauthorized, map[string]string{ "error": "device not authenticated", }) } if !m.hasPermission(device.DeviceType, permission) { return c.JSON(http.StatusForbidden, map[string]string{ "error": "insufficient permissions", }) } return next(c) } } } func (m *DeviceAuthMiddleware) hasPermission(deviceType string, permission string) bool { permissions := map[string][]string{ "koreader": {"sync:progress", "sync:annotations", "sync:metadata"}, "kobo": {"sync:progress", "sync:annotations", "sync:metadata"}, "web": {"sync:progress", "sync:annotations", "sync:metadata", "device:manage"}, "mobile": {"sync:progress", "sync:annotations", "sync:metadata"}, } devicePerms, exists := permissions[deviceType] if !exists { return false } for _, p := range devicePerms { if p == permission { return true } } return false } func (m *DeviceAuthMiddleware) UpdateLastSeen(next echo.HandlerFunc) echo.HandlerFunc { return func(c echo.Context) error { err := next(c) deviceIDBytes, ok := c.Get("device_id").([16]byte) if ok { pgDeviceID := pgtype.UUID{Bytes: deviceIDBytes, Valid: true} m.db.UpdateDeviceLastSeen(c.Request().Context(), pgDeviceID) } return err } } func (m *DeviceAuthMiddleware) ValidateDeviceToken(token string) (database.Devices, error) { return m.db.GetDeviceByAuthToken(context.Background(), token) }