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/v5" ) 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 settings *database.SettingsRegistry } func NewDeviceAuthMiddleware(db *database.Queries) *DeviceAuthMiddleware { return &DeviceAuthMiddleware{ db: db, rateLimiter: NewDeviceRateLimiter(), } } // SetSettings wires the tunable settings registry so device rate limits are // read live on each authenticated request. func (m *DeviceAuthMiddleware) SetSettings(s *database.SettingsRegistry) { m.settings = s } // rateLimitConfig returns the active device rate limits from the registry, or // the historical defaults when no registry is wired. func (m *DeviceAuthMiddleware) rateLimitConfig() DeviceRateLimitConfig { if m.settings != nil { dl := m.settings.DeviceRateLimits() return DeviceRateLimitConfig{ SyncRequestsPerMinute: dl.Sync, ProgressUpdatesPerMinute: dl.Progress, MetadataRequestsPerMinute: dl.Metadata, } } return DeviceRateLimitConfig{ SyncRequestsPerMinute: DefaultSyncRequestsPerMinute, ProgressUpdatesPerMinute: DefaultProgressUpdatesPerMinute, MetadataRequestsPerMinute: DefaultMetadataRequestsPerMinute, } } // rateLimitForRequestType returns the configured per-minute limit for a given // request type, for use in X-RateLimit-* headers. func (m *DeviceAuthMiddleware) rateLimitForRequestType(requestType string, config DeviceRateLimitConfig) int { switch requestType { case "progress": return config.ProgressUpdatesPerMinute case "metadata": return config.MetadataRequestsPerMinute default: // "sync" and any unknown type return config.SyncRequestsPerMinute } } func (m *DeviceAuthMiddleware) Authenticate(next echo.HandlerFunc) echo.HandlerFunc { return func(c *echo.Context) error { var device database.Devices var err error var token string var urlToken string var queryToken string // For OPDS routes, validate deviceId is a valid UUID before authentication // This allows returning 400 Bad Request for invalid UUIDs instead of 401 if strings.HasPrefix(c.Request().URL.Path, "/opds/devices/") { deviceID := c.Param("deviceId") if deviceID != "" { if _, err := uuid.Parse(deviceID); err != nil { return c.JSON(http.StatusBadRequest, map[string]string{"error": "invalid device ID"}) } } // Also validate bookId for download, cover, and formats endpoints bookID := c.Param("bookId") if bookID != "" { if _, err := uuid.Parse(bookID); err != nil { return c.JSON(http.StatusBadRequest, map[string]string{"error": "invalid book ID"}) } } } // Method 1: Try Bearer token header (KOReader, API clients, OPDS) authHeader := c.Request().Header.Get("Authorization") if authHeader != "" { 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 { // Found device via Bearer token, continue to validation goto validateDevice } } // Method 2: Try URL path parameter (Kobo sync, OPDS) // Route format: /api/sync/kobo/:token/... urlToken = c.Param("token") if urlToken != "" { device, err = m.db.GetDeviceByAuthToken(c.Request().Context(), urlToken) if err == nil { // Found device via URL path token, continue to validation goto validateDevice } } // Method 3: Try query parameter (OPDS catalog access) // URL format: /opds/devices/:deviceId/catalog?token=... queryToken = c.QueryParam("token") if queryToken != "" { device, err = m.db.GetDeviceByAuthToken(c.Request().Context(), queryToken) if err == nil { // Found device via query parameter token, continue to validation goto validateDevice } } // All authentication methods failed return c.JSON(http.StatusUnauthorized, map[string]string{ "error": "authentication required - use Bearer token or API key", }) validateDevice: 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 := m.rateLimitConfig() limitForType := m.rateLimitForRequestType(requestType, config) if !m.rateLimiter.CheckRateLimit(deviceID, requestType, config) { remaining := m.rateLimiter.GetRemainingRequests(deviceID, requestType, config) c.Response().Header().Set("X-RateLimit-Limit", strconv.Itoa(limitForType)) 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", strconv.Itoa(limitForType)) 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) }