Files
dumpsterChat/internal/gateway/client.go
T

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