Files
dumpsterChat/internal/gateway/client.go
T
hobokenchicken 02db6c719c fix(gateway): hash session tokens in WS auth; fix WS reconnect backoff
The token-hashing commit (57aec2c) never updated ServeWS to hash tokens
before querying the sessions table. Cookie and message-frame auth both
compared raw tokens against stored hashes, so every WS connection failed.

Also moved reconnectDelay to module scope so the exponential backoff
survives across connect() calls, and resets on successful open.
2026-07-27 13:22:43 -04:00

524 lines
14 KiB
Go

package gateway
import (
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"encoding/json"
"log/slog"
"net/http"
"strings"
"time"
"github.com/gorilla/websocket"
)
const (
writeWait = 10 * time.Second
pongWait = 60 * time.Second
pingPeriod = (pongWait * 9) / 10
maxMessageSize = 4096
)
var upgrader = websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
CheckOrigin: func(r *http.Request) bool { return true },
}
// SetAllowedOrigins replaces the default upgrader with one that validates
// the Origin header. Accepts if:
// 1. Exact match in origins list, OR
// 2. Hostname matches a trusted hostname, OR
// 3. Hostname matches the request's own Host header (same-origin), OR
// 4. Localhost variant
func SetAllowedOrigins(origins []string) {
originSet := make(map[string]struct{}, len(origins))
hostSet := make(map[string]struct{}, len(origins))
for _, o := range origins {
originSet[o] = struct{}{}
h := o
if idx := strings.Index(h, "://"); idx >= 0 {
h = h[idx+3:]
}
if idx := strings.Index(h, "/"); idx >= 0 {
h = h[:idx]
}
if idx := strings.LastIndex(h, ":"); idx >= 0 {
h = h[:idx]
}
if h != "" {
hostSet[h] = struct{}{}
}
}
upgrader = websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
CheckOrigin: func(r *http.Request) bool {
origin := r.Header.Get("Origin")
if origin == "" {
return true
}
if _, ok := originSet[origin]; ok {
return true
}
oh := extractHostname(origin)
// Match trusted hostnames
if _, ok := hostSet[oh]; ok {
return true
}
// Match request's own Host header (same-origin)
reqHost := extractHostname(r.Host)
if reqHost != "" && oh == reqHost {
return true
}
// Localhost variants
if oh == "localhost" || oh == "127.0.0.1" || oh == "::1" {
return true
}
return false
},
}
}
func extractHostname(s string) string {
h := s
if idx := strings.Index(h, "://"); idx >= 0 {
h = h[idx+3:]
}
if idx := strings.Index(h, "/"); idx >= 0 {
h = h[:idx]
}
if idx := strings.LastIndex(h, ":"); idx >= 0 {
h = h[:idx]
}
h = strings.TrimPrefix(h, "[")
h = strings.TrimSuffix(h, "]")
return h
}
// Client is a middleman between the websocket connection and the hub.
type Client struct {
Hub *Hub
Conn *websocket.Conn
UserID string
Username string
IsBot bool
BotID string
BotName string
send chan []byte
}
// readPump pumps messages from the websocket connection to the hub.
func (c *Client) readPump() {
defer func() {
c.Hub.Unregister(c)
c.Conn.Close()
}()
c.Conn.SetReadLimit(maxMessageSize)
c.Conn.SetReadDeadline(time.Now().Add(pongWait))
c.Conn.SetPongHandler(func(string) error {
c.Conn.SetReadDeadline(time.Now().Add(pongWait))
return nil
})
for {
_, raw, err := c.Conn.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
c.Hub.logger.Warn("websocket read error", "error", err, "user_id", c.UserID)
}
break
}
var event Event
if err := json.Unmarshal(raw, &event); err != nil {
c.Hub.logger.Warn("invalid event JSON", "error", err, "user_id", c.UserID)
continue
}
c.Hub.RecordActivity(c.UserID)
switch event.Type {
case EventTypingStart:
// Enrich with sender info so receivers know who is typing
var data map[string]interface{}
if raw, ok := event.Data.(json.RawMessage); ok {
json.Unmarshal(raw, &data)
}
if data == nil {
data = make(map[string]interface{})
}
data["user_id"] = c.UserID
data["username"] = c.Username
c.Hub.BroadcastEvent(Event{Type: event.Type, Data: data})
case EventPresenceUpdate:
c.Hub.BroadcastEvent(event)
case EventVoiceJoin, EventVoiceLeave, EventVoiceMute, EventVoiceDeafen:
// Broadcast voice state changes to all clients with sender info
var vData map[string]interface{}
if raw, ok := event.Data.(json.RawMessage); ok {
json.Unmarshal(raw, &vData)
}
if vData == nil {
vData = make(map[string]interface{})
}
vData["user_id"] = c.UserID
vData["username"] = c.Username
c.Hub.BroadcastEvent(Event{Type: event.Type, Data: vData})
case EventVoiceWhisper:
// Forward voice whispers only to the target user, not broadcast
var whisperData struct {
TargetUserID string `json:"target_user_id"`
FromUserID string `json:"from_user_id"`
FromUsername string `json:"from_username"`
Message string `json:"message,omitempty"`
}
if raw, ok := event.Data.(json.RawMessage); ok {
if err := json.Unmarshal(raw, &whisperData); err == nil {
whisperData.FromUserID = c.UserID
c.Hub.SendToUser(whisperData.TargetUserID, Event{
Type: EventVoiceWhisper,
Data: whisperData,
})
}
}
case BotSendMessage:
if !c.IsBot {
c.Hub.logger.Warn("non-bot client sent SEND_MESSAGE", "user_id", c.UserID)
continue
}
c.handleBotSendMessage(event.Data)
case BotDeleteMessage:
if !c.IsBot {
c.Hub.logger.Warn("non-bot client sent DELETE_MESSAGE", "user_id", c.UserID)
continue
}
c.handleBotDeleteMessage(event.Data)
default:
c.Hub.logger.Info("received event from client", "type", event.Type, "user_id", c.UserID)
}
}
}
// writePump pumps messages from the hub to the websocket connection.
func (c *Client) writePump() {
ticker := time.NewTicker(pingPeriod)
defer func() {
ticker.Stop()
c.Conn.Close()
}()
for {
select {
case message, ok := <-c.send:
c.Conn.SetWriteDeadline(time.Now().Add(writeWait))
if !ok {
c.Conn.WriteMessage(websocket.CloseMessage, []byte{})
return
}
w, err := c.Conn.NextWriter(websocket.TextMessage)
if err != nil {
return
}
w.Write(message)
n := len(c.send)
for i := 0; i < n; i++ {
w.Write([]byte{'\n'})
w.Write(<-c.send)
}
if err := w.Close(); err != nil {
return
}
case <-ticker.C:
c.Conn.SetWriteDeadline(time.Now().Add(writeWait))
if err := c.Conn.WriteMessage(websocket.PingMessage, nil); err != nil {
return
}
}
}
}
// authMessage is the first frame the client must send after connecting (for TUI/bot clients).
type authMessage struct {
Token string `json:"token"`
}
// ServeWS handles websocket requests from the peer.
// It authenticates via cookie (browser) or message-frame (TUI/bots).
func ServeWS(db *sql.DB, hub *Hub, logger *slog.Logger, w http.ResponseWriter, r *http.Request, cookieName string) {
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
logger.Error("websocket upgrade failed", "error", err)
return
}
var userID, username string
// Collect candidate tokens from Query param, Authorization header, and Cookies.
var candidateTokens []string
if qToken := r.URL.Query().Get("token"); qToken != "" {
candidateTokens = append(candidateTokens, qToken)
}
authHeader := r.Header.Get("Authorization")
if len(authHeader) > 7 && authHeader[:7] == "Bearer " {
candidateTokens = append(candidateTokens, authHeader[7:])
}
for _, c := range r.Cookies() {
if c.Name == cookieName && c.Value != "" {
candidateTokens = append(candidateTokens, c.Value)
}
}
// Test candidates against Postgres sessions (tokens are stored hashed)
for _, token := range candidateTokens {
err := db.QueryRowContext(context.Background(),
`SELECT u.id, u.username FROM sessions s JOIN users u ON u.id = s.user_id WHERE s.token = $1 AND s.expires_at > NOW()`,
hashToken(token),
).Scan(&userID, &username)
if err == nil && userID != "" {
break
}
}
// Fall back to message-frame auth if no candidate token authenticated (wait up to 5s)
if userID == "" {
conn.SetReadDeadline(time.Now().Add(5 * time.Second))
_, raw, readErr := conn.ReadMessage()
if readErr == nil {
var auth authMessage
if jsonErr := json.Unmarshal(raw, &auth); jsonErr == nil && auth.Token != "" {
_ = db.QueryRowContext(context.Background(),
`SELECT u.id, u.username FROM sessions s JOIN users u ON u.id = s.user_id WHERE s.token = $1 AND s.expires_at > NOW()`,
hashToken(auth.Token),
).Scan(&userID, &username)
}
}
}
if userID == "" {
logger.Warn("ws auth: no valid session found")
conn.WriteMessage(websocket.TextMessage, []byte(`{"type":"error","error":"invalid_session"}`))
conn.Close()
return
}
conn.SetReadDeadline(time.Time{})
conn.WriteMessage(websocket.TextMessage, []byte(`{"type":"ready"}`))
client := &Client{
Hub: hub,
Conn: conn,
UserID: userID,
Username: username,
send: make(chan []byte, 256),
}
hub.Register(client)
go client.writePump()
go client.readPump()
}
// ServeBotWS handles websocket requests from bot clients.
// Authenticates via ?token= query param (bot token, hashed lookup).
func ServeBotWS(db *sql.DB, hub *Hub, logger *slog.Logger, w http.ResponseWriter, r *http.Request) {
token := r.URL.Query().Get("token")
if token == "" {
http.Error(w, `{"error":"token query param required"}`, http.StatusBadRequest)
return
}
// Hash the token and look up the bot.
tokenHash := hashToken(token)
var botID, botName, ownerID string
err := db.QueryRowContext(r.Context(),
`SELECT id, name, owner_id FROM bots WHERE token = $1`,
tokenHash,
).Scan(&botID, &botName, &ownerID)
if err != nil {
logger.Warn("bot ws auth: invalid token")
http.Error(w, `{"error":"invalid bot token"}`, http.StatusUnauthorized)
return
}
// Verify the bot is added to at least one server.
var serverCount int
err = db.QueryRowContext(r.Context(),
`SELECT COUNT(*) FROM bot_servers WHERE bot_id = $1`, botID,
).Scan(&serverCount)
if err != nil || serverCount == 0 {
logger.Warn("bot ws auth: bot not added to any server", "bot_id", botID)
http.Error(w, `{"error":"bot not added to any server"}`, http.StatusForbidden)
return
}
conn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
logger.Error("bot websocket upgrade failed", "error", err)
return
}
// Load bot server memberships into hub so BroadcastToServer works.
hub.RefreshUserServers(ownerID)
conn.SetReadDeadline(time.Time{})
conn.WriteMessage(websocket.TextMessage, []byte(`{"type":"ready","bot_id":"`+botID+`"}`))
client := &Client{
Hub: hub,
Conn: conn,
UserID: ownerID,
Username: botName,
IsBot: true,
BotID: botID,
BotName: botName,
send: make(chan []byte, 256),
}
hub.Register(client)
go client.writePump()
go client.readPump()
}
// hashToken returns the SHA-256 hex digest of a token.
func hashToken(token string) string {
h := sha256.Sum256([]byte(token))
return hex.EncodeToString(h[:])
}
// handleBotSendMessage processes a SEND_MESSAGE action from a bot client.
func (c *Client) handleBotSendMessage(data interface{}) {
var payload struct {
ChannelID string `json:"channel_id"`
Content string `json:"content"`
}
raw, ok := data.(json.RawMessage)
if !ok {
return
}
if err := json.Unmarshal(raw, &payload); err != nil || payload.ChannelID == "" || payload.Content == "" {
c.Hub.logger.Warn("bot SEND_MESSAGE: invalid payload")
return
}
if len(payload.Content) > 4000 {
payload.Content = payload.Content[:4000]
}
// Verify the bot is in the server that owns this channel.
serverID, err := c.Hub.ServerIDForChannel(context.Background(), payload.ChannelID)
if err != nil {
c.Hub.logger.Warn("bot SEND_MESSAGE: channel not found", "channel_id", payload.ChannelID)
return
}
var inServer bool
err = c.Hub.db.QueryRowContext(context.Background(),
`SELECT EXISTS(SELECT 1 FROM bot_servers WHERE bot_id = $1 AND server_id = $2)`,
c.BotID, serverID,
).Scan(&inServer)
if err != nil || !inServer {
c.Hub.logger.Warn("bot SEND_MESSAGE: bot not in server", "bot_id", c.BotID, "server_id", serverID)
return
}
// Insert the message with bot_id set.
var msgID, createdAt string
err = c.Hub.db.QueryRowContext(context.Background(),
`INSERT INTO messages (channel_id, author_id, content, bot_id)
VALUES ($1, $2, $3, $4)
RETURNING id, created_at::text`,
payload.ChannelID, c.UserID, payload.Content, c.BotID,
).Scan(&msgID, &createdAt)
if err != nil {
c.Hub.logger.Error("bot SEND_MESSAGE: insert failed", "error", err)
return
}
// Broadcast MESSAGE_CREATE to the server.
c.Hub.BroadcastToServer(serverID, Event{
Type: EventMessageCreate,
Data: map[string]interface{}{
"id": msgID,
"channel_id": payload.ChannelID,
"author_id": c.UserID,
"author_username": c.BotName,
"author_display_name": nil,
"author_bot": true,
"bot_id": c.BotID,
"bot_name": c.BotName,
"content": payload.Content,
"reply_to": nil,
"edited_at": nil,
"pinned": false,
"created_at": createdAt,
"embeds": []interface{}{},
"reactions": []interface{}{},
},
})
}
// handleBotDeleteMessage processes a DELETE_MESSAGE action from a bot client.
func (c *Client) handleBotDeleteMessage(data interface{}) {
var payload struct {
ChannelID string `json:"channel_id"`
MessageID string `json:"message_id"`
}
raw, ok := data.(json.RawMessage)
if !ok {
return
}
if err := json.Unmarshal(raw, &payload); err != nil || payload.ChannelID == "" || payload.MessageID == "" {
c.Hub.logger.Warn("bot DELETE_MESSAGE: invalid payload")
return
}
// Verify the bot is in the server that owns this channel.
serverID, err := c.Hub.ServerIDForChannel(context.Background(), payload.ChannelID)
if err != nil {
return
}
var inServer bool
err = c.Hub.db.QueryRowContext(context.Background(),
`SELECT EXISTS(SELECT 1 FROM bot_servers WHERE bot_id = $1 AND server_id = $2)`,
c.BotID, serverID,
).Scan(&inServer)
if err != nil || !inServer {
return
}
// Delete the message (only if it exists in this channel).
result, err := c.Hub.db.ExecContext(context.Background(),
`DELETE FROM messages WHERE id = $1 AND channel_id = $2`,
payload.MessageID, payload.ChannelID,
)
if err != nil {
c.Hub.logger.Error("bot DELETE_MESSAGE: delete failed", "error", err)
return
}
rows, _ := result.RowsAffected()
if rows == 0 {
return
}
// Broadcast MESSAGE_DELETE.
c.Hub.BroadcastToServer(serverID, Event{
Type: EventMessageDelete,
Data: map[string]string{
"id": payload.MessageID,
"channel_id": payload.ChannelID,
},
})
}