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.
This commit is contained in:
@@ -0,0 +1,230 @@
|
||||
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(),
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user