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, }, }) }