Files
bookhoard/internal/handlers/websocket.go
T
john-okeefe 317a4e82a0 feat: add WebSocket handler with dual authentication
Add WSHandler for WebSocket connection management:
- Upgrade HTTP to WebSocket connections
- Dual authentication support:
  - JWT token via query parameter (web clients)
  - Bearer token via Authorization header (devices)
- Client info extraction for users and devices
- Separate read and write pumps for concurrent I/O
- Ping/pong heartbeat mechanism (30s interval)
- Initial state delivery on connection
- Connection cleanup on disconnect

Implements Week 9 WebSocket endpoint functionality from
Universal Sync Implementation Guide.
2026-01-30 21:47:11 -05:00

231 lines
6.2 KiB
Go

package handlers
import (
"bookmann/internal/database"
"bookmann/internal/middleware"
"bookmann/internal/sync"
"context"
"log"
"net/http"
"time"
jwt "github.com/golang-jwt/jwt/v5"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/jackc/pgx/v5/pgtype"
"github.com/labstack/echo/v4"
)
var upgrader = websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
CheckOrigin: func(r *http.Request) bool {
return true
},
}
type WSHandler struct {
db *database.Queries
connManager *sync.ConnectionManager
jwtSecret string
deviceAuthMiddleware *middleware.DeviceAuthMiddleware
}
func NewWSHandler(db *database.Queries, connManager *sync.ConnectionManager, jwtSecret string, deviceAuth *middleware.DeviceAuthMiddleware) *WSHandler {
return &WSHandler{
db: db,
connManager: connManager,
jwtSecret: jwtSecret,
deviceAuthMiddleware: deviceAuth,
}
}
type WSMessage struct {
Type string `json:"type"`
Timestamp string `json:"timestamp"`
Data map[string]interface{} `json:"data"`
}
type ClientInfo struct {
UserID string
DeviceID string
DeviceType string
DeviceName string
IsDevice bool
}
func (h *WSHandler) HandleWebSocket(c echo.Context) error {
token := c.QueryParam("token")
if token == "" {
return c.JSON(http.StatusUnauthorized, map[string]string{"error": "token required"})
}
clientInfo, err := h.authenticateClient(token, c.Request().Header.Get("Authorization"))
if err != nil {
return c.JSON(http.StatusUnauthorized, map[string]string{"error": err.Error()})
}
ws, err := upgrader.Upgrade(c.Response(), c.Request(), nil)
if err != nil {
log.Printf("WebSocket upgrade failed: %v", err)
return err
}
conn := h.createConnection(clientInfo)
h.connManager.AddConnection(conn)
go h.readPump(conn, ws, clientInfo)
go h.writePump(conn, ws, clientInfo)
return nil
}
func (h *WSHandler) authenticateClient(token string, authHeader string) (*ClientInfo, error) {
if authHeader != "" && len(authHeader) > 7 && authHeader[:7] == "Bearer " {
deviceToken := authHeader[7:]
device, err := h.deviceAuthMiddleware.ValidateDeviceToken(deviceToken)
if err == nil {
userID := uuid.UUID(device.UserID.Bytes).String()
return &ClientInfo{
UserID: userID,
DeviceID: uuid.UUID(device.ID.Bytes).String(),
DeviceType: device.DeviceType,
DeviceName: device.DeviceName,
IsDevice: true,
}, nil
}
}
parsedToken, err := jwt.Parse(token, func(token *jwt.Token) (interface{}, error) {
return []byte(h.jwtSecret), nil
})
if err == nil && parsedToken.Valid {
claims := parsedToken.Claims.(jwt.MapClaims)
userID := claims["user_id"].(string)
return &ClientInfo{
UserID: userID,
DeviceID: "web-" + userID,
DeviceType: "web",
DeviceName: "Web Client",
IsDevice: false,
}, nil
}
return nil, echo.NewHTTPError(http.StatusUnauthorized, "invalid token")
}
func (h *WSHandler) createConnection(clientInfo *ClientInfo) *sync.DeviceConnection {
return &sync.DeviceConnection{
ID: clientInfo.DeviceID,
UserID: clientInfo.UserID,
DeviceType: clientInfo.DeviceType,
DeviceName: clientInfo.DeviceName,
Connected: time.Now(),
LastPing: time.Now(),
Send: make(chan sync.BroadcastMessage, 100),
Disconnected: make(chan struct{}),
}
}
func (h *WSHandler) readPump(conn *sync.DeviceConnection, ws *websocket.Conn, clientInfo *ClientInfo) {
defer func() {
h.connManager.RemoveConnection(conn.ID)
ws.Close()
}()
ws.SetReadDeadline(time.Now().Add(90 * time.Second))
ws.SetPongHandler(func(string) error {
ws.SetReadDeadline(time.Now().Add(90 * time.Second))
conn.LastPing = time.Now()
return nil
})
for {
_, message, err := ws.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseAbnormalClosure) {
log.Printf("WebSocket error for %s: %v", conn.DeviceName, err)
}
break
}
h.handleMessage(message, conn, clientInfo)
}
}
func (h *WSHandler) writePump(conn *sync.DeviceConnection, ws *websocket.Conn, clientInfo *ClientInfo) {
ticker := time.NewTicker(30 * time.Second)
defer func() {
ticker.Stop()
ws.Close()
}()
initialState := h.getInitialState(clientInfo.UserID)
conn.Send <- sync.BroadcastMessage{
Type: sync.MessageTypeInitial,
Timestamp: time.Now().Format(time.RFC3339),
Data: initialState,
}
for {
select {
case msg, ok := <-conn.Send:
if !ok {
return
}
ws.SetWriteDeadline(time.Now().Add(10 * time.Second))
if err := ws.WriteJSON(msg); err != nil {
log.Printf("WebSocket write error for %s: %v", conn.DeviceName, err)
return
}
case <-conn.Disconnected:
return
case <-ticker.C:
ws.SetWriteDeadline(time.Now().Add(10 * time.Second))
if err := ws.WriteMessage(websocket.PingMessage, nil); err != nil {
return
}
}
}
}
func (h *WSHandler) handleMessage(message []byte, conn *sync.DeviceConnection, clientInfo *ClientInfo) {
log.Printf("WebSocket message from %s: %s", conn.DeviceName, string(message))
}
func (h *WSHandler) getInitialState(userID string) map[string]interface{} {
userUUID, err := uuid.Parse(userID)
if err != nil {
return map[string]interface{}{"error": "invalid user ID"}
}
pgUserID := pgtype.UUID{Bytes: userUUID, Valid: true}
mediaItems, err := h.db.GetUserMediaItemsForSync(context.Background(), pgUserID)
if err != nil {
return map[string]interface{}{"error": "failed to fetch initial state"}
}
progressMap := make(map[string]map[string]interface{})
for _, item := range mediaItems {
progress, err := h.db.GetUniversalProgress(context.Background(), database.GetUniversalProgressParams{
MediaItemID: item.ID,
UserID: pgUserID,
})
if err == nil {
bookID := uuid.UUID(item.ID.Bytes).String()
progressMap[bookID] = map[string]interface{}{
"percentage": progress.Percentage.Float64,
"current_page": progress.CurrentPage.Int32,
"total_pages": progress.TotalPages.Int32,
"last_read": progress.LastReadAt.Time,
}
}
}
return map[string]interface{}{
"progress": progressMap,
"devices": h.connManager.GetConnectionStats(),
}
}