live-sfu-demo/internal/server/chat.go

295 lines
8.2 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package server
import (
"encoding/json"
"fmt"
"log"
"net"
"net/http"
"strings"
"sync"
"time"
)
// chatMessage 是房间弹幕/聊天的一条广播消息。
type chatMessage struct {
ID string `json:"id"`
Room string `json:"room"`
User string `json:"user"`
Role string `json:"role"`
Message string `json:"message"`
Color string `json:"color,omitempty"`
Host bool `json:"host,omitempty"`
TS int64 `json:"ts"`
}
// chatHub 维护房间 -> 弹幕订阅者的内存广播(与 roomHub 同构,重启不恢复历史)。
type chatHub struct {
mu sync.RWMutex
rooms map[string]map[chan []byte]struct{}
// recentMsgs 保存每个房间最近若干条弹幕,用于新订阅者进场回放(重启不恢复,符合内存态设计)。
recentMsgs map[string][]*chatMessage
}
// recentLimit 是单房间保留的弹幕历史条数上限。
const recentLimit = 50
func newChatHub() *chatHub {
return &chatHub{rooms: map[string]map[chan []byte]struct{}{}, recentMsgs: map[string][]*chatMessage{}}
}
func (c *chatHub) subscribe(room string) (chan []byte, func()) {
c.mu.Lock()
defer c.mu.Unlock()
ch := make(chan []byte, 16)
if c.rooms[room] == nil {
c.rooms[room] = map[chan []byte]struct{}{}
}
if c.recentMsgs[room] == nil {
c.recentMsgs[room] = nil
}
c.rooms[room][ch] = struct{}{}
return ch, func() {
c.mu.Lock()
defer c.mu.Unlock()
if m, ok := c.rooms[room]; ok {
delete(m, ch)
if len(m) == 0 {
delete(c.rooms, room)
// 房间再无订阅者时清空历史,避免内存随房间数无界增长。
delete(c.recentMsgs, room)
}
}
}
}
// recent 返回房间最近若干条弹幕的快照(副本,调用方无需持有锁)。
func (c *chatHub) recent(room string) []*chatMessage {
c.mu.RLock()
defer c.mu.RUnlock()
out := make([]*chatMessage, len(c.recentMsgs[room]))
copy(out, c.recentMsgs[room])
return out
}
func (c *chatHub) broadcast(room string, msg *chatMessage) {
data, err := json.Marshal(msg)
if err != nil {
return
}
c.mu.Lock()
c.recentMsgs[room] = append(c.recentMsgs[room], msg)
if n := len(c.recentMsgs[room]); n > recentLimit {
c.recentMsgs[room] = c.recentMsgs[room][n-recentLimit:]
}
for ch := range c.rooms[room] {
select {
case ch <- data:
default:
}
}
c.mu.Unlock()
}
// chatRateLimiter 是弹幕发送的滑动窗口限流(按 房间 + 发送者 + IP 维度),用于防刷屏。
// 纯内存实现,重启即清空,与弹幕本身的内存态一致;后台定期清理过期 key 防止 map 无界增长。
type chatRateLimiter struct {
mu sync.Mutex
win map[string][]time.Time
limit int
window time.Duration
}
func newChatRateLimiter(limit int, window time.Duration) *chatRateLimiter {
r := &chatRateLimiter{win: map[string][]time.Time{}, limit: limit, window: window}
go r.cleanup()
return r
}
// allow 在窗口内计数未超上限时记入并返回 true,否则返回 false(应拒绝本次发送)。
func (r *chatRateLimiter) allow(key string, now time.Time) bool {
r.mu.Lock()
defer r.mu.Unlock()
cut := now.Add(-r.window)
ts := r.win[key]
kept := ts[:0]
for _, t := range ts {
if t.After(cut) {
kept = append(kept, t)
}
}
if len(kept) >= r.limit {
r.win[key] = kept
return false
}
r.win[key] = append(kept, now)
return true
}
// cleanup 周期性删除已无活跃计数的 key。
func (r *chatRateLimiter) cleanup() {
ticker := time.NewTicker(time.Minute)
defer ticker.Stop()
for range ticker.C {
r.mu.Lock()
cut := time.Now().Add(-r.window)
for k, ts := range r.win {
if len(ts) == 0 || ts[len(ts)-1].Before(cut) {
delete(r.win, k)
}
}
r.mu.Unlock()
}
}
// clientIP 从代理头或 RemoteAddr 取客户端 IP(取 X-Forwarded-For 第一段)。
func clientIP(r *http.Request) string {
if x := r.Header.Get("X-Forwarded-For"); x != "" {
if i := strings.IndexByte(x, ','); i >= 0 {
return strings.TrimSpace(x[:i])
}
return strings.TrimSpace(x)
}
if x := r.Header.Get("X-Real-IP"); x != "" {
return strings.TrimSpace(x)
}
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return host
}
// handleRoomChat 以 SSE 推送房间弹幕(前端据此渲染叠加层 / 聊天列表)。
func (s *Server) handleRoomChat(w http.ResponseWriter, r *http.Request) {
room := r.PathValue("room")
flusher, ok := w.(http.Flusher)
if !ok {
http.Error(w, "streaming unsupported", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
ch, unsub := s.chat.subscribe(room)
defer unsub()
// 进场回放最近若干条弹幕,避免新订阅者面对空房间(非阻塞,填满通道即止)。
for _, m := range s.chat.recent(room) {
b, err := json.Marshal(m)
if err != nil {
continue
}
select {
case ch <- b:
default:
break
}
}
ticker := time.NewTicker(25 * time.Second)
defer ticker.Stop()
for {
select {
case <-r.Context().Done():
return
case msg := <-ch:
fmt.Fprintf(w, "event: chat\ndata: %s\n\n", msg)
flusher.Flush()
case <-ticker.C:
fmt.Fprintf(w, ": ping\n\n")
flusher.Flush()
}
}
}
// handleRoomChatPost 接收一条弹幕并广播给房间内所有订阅者。
func (s *Server) handleRoomChatPost(w http.ResponseWriter, r *http.Request) {
room := r.PathValue("room")
var body struct {
Message string `json:"message"`
Color string `json:"color"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil || body.Message == "" {
http.Error(w, "bad request", http.StatusBadRequest)
return
}
if len(body.Message) > 500 {
body.Message = body.Message[:500]
}
// 限流:按 房间 + 发送者 + IP 维度做滑动窗口,防刷屏。
user, role, id := "guest", "guest", "guest"
if s.auth != nil {
if _, u, err := s.auth.VerifyToken(extractAuthToken(r)); err == nil {
user, role, id = u.Username, u.Role, u.Username
}
}
if s.chatRL != nil {
key := room + "/" + id + "/" + clientIP(r)
if !s.chatRL.allow(key, time.Now()) {
http.Error(w, "too many messages, please slow down", http.StatusTooManyRequests)
return
}
}
msg := &chatMessage{
ID: randomID(),
Room: room,
User: user,
Role: role,
Message: body.Message,
Color: body.Color,
TS: time.Now().Unix(),
}
s.chat.broadcast(room, msg)
log.Printf("[chat] room=%s user=%s role=%s: %s", room, user, role, body.Message)
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]any{"ok": true, "id": msg.ID})
}
// handleRoomBroadcast 是主播的独立外部推送接口:仅 publisher/admin 可调用(authWrap 已校验 room:publish),
// 推送的弹幕标记 host=true,经同一 chatHub 广播,主播端与观众端实时可见。供 OBS 脚本 / 机器人 / 管理工具使用。
func (s *Server) handleRoomBroadcast(w http.ResponseWriter, r *http.Request) {
room := r.PathValue("room")
var body struct {
Message string `json:"message"`
Color string `json:"color"`
}
if err := json.NewDecoder(r.Body).Decode(&body); err != nil || body.Message == "" {
http.Error(w, "bad request", http.StatusBadRequest)
return
}
if len(body.Message) > 500 {
body.Message = body.Message[:500]
}
if s.chat == nil {
http.Error(w, "chat not available", http.StatusServiceUnavailable)
return
}
// 身份来自 token(外部工具用 Authorization: Bearer <token> 调用)
user, role, id := "guest", "guest", "guest"
if s.auth != nil {
if _, u, err := s.auth.VerifyToken(extractAuthToken(r)); err == nil {
user, role, id = u.Username, u.Role, u.Username
}
}
if s.chatRL != nil {
key := room + "/broadcast/" + id + "/" + clientIP(r)
if !s.chatRL.allow(key, time.Now()) {
http.Error(w, "too many messages, please slow down", http.StatusTooManyRequests)
return
}
}
msg := &chatMessage{
ID: randomID(),
Room: room,
User: user,
Role: role,
Message: body.Message,
Color: body.Color,
Host: true,
TS: time.Now().Unix(),
}
s.chat.broadcast(room, msg)
log.Printf("[chat:broadcast] room=%s user=%s role=%s: %s", room, user, role, body.Message)
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]any{"ok": true, "id": msg.ID})
}