524 lines
14 KiB
Go
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,
|
|
},
|
|
})
|
|
}
|