fix(relay): harden lifecycle and protocol handling
This commit is contained in:
+245
-226
@@ -25,7 +25,6 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
maxRoomSize = 8
|
||||
rateBurst = 30
|
||||
rateSustained = 10
|
||||
cleanupInterval = 5 * time.Minute
|
||||
@@ -34,7 +33,6 @@ const (
|
||||
writeWait = 10 * time.Second
|
||||
pongWait = 60 * time.Second
|
||||
pingInterval = 30 * time.Second
|
||||
maxMessageSize = 64 * 1024
|
||||
maxLogSize = 1 * 1024 * 1024 // 1MB
|
||||
logMaxAge = 3 * 24 * time.Hour
|
||||
logIDLength = 5
|
||||
@@ -62,142 +60,6 @@ var upgrader = websocket.Upgrader{
|
||||
CheckOrigin: func(r *http.Request) bool { return true },
|
||||
}
|
||||
|
||||
// --- Rate limiter (token bucket) ---
|
||||
|
||||
type rateLimiter struct {
|
||||
tokens float64
|
||||
maxTokens float64
|
||||
refillRate float64
|
||||
lastTime time.Time
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func newRateLimiter(burst, sustained int) *rateLimiter {
|
||||
return &rateLimiter{
|
||||
tokens: float64(burst),
|
||||
maxTokens: float64(burst),
|
||||
refillRate: float64(sustained),
|
||||
lastTime: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
func (rl *rateLimiter) allow() bool {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
elapsed := now.Sub(rl.lastTime).Seconds()
|
||||
rl.lastTime = now
|
||||
|
||||
rl.tokens += elapsed * rl.refillRate
|
||||
if rl.tokens > rl.maxTokens {
|
||||
rl.tokens = rl.maxTokens
|
||||
}
|
||||
|
||||
if rl.tokens < 1 {
|
||||
return false
|
||||
}
|
||||
rl.tokens--
|
||||
return true
|
||||
}
|
||||
|
||||
// stale reports whether a limiter hasn't been touched in over 10 minutes —
|
||||
// safe to GC from a per-IP map.
|
||||
func (rl *rateLimiter) stale(now time.Time) bool {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
return now.Sub(rl.lastTime) > 10*time.Minute
|
||||
}
|
||||
|
||||
// --- Connection tracker (per-IP limits) ---
|
||||
|
||||
type connTracker struct {
|
||||
mu sync.Mutex
|
||||
perIP map[string]int
|
||||
ipRate map[string]*rateLimiter
|
||||
roomsPerIP map[string]int
|
||||
globalCount int
|
||||
}
|
||||
|
||||
func newConnTracker() *connTracker {
|
||||
return &connTracker{
|
||||
perIP: make(map[string]int),
|
||||
ipRate: make(map[string]*rateLimiter),
|
||||
roomsPerIP: make(map[string]int),
|
||||
}
|
||||
}
|
||||
|
||||
func (ct *connTracker) tryConnect(ip string) bool {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
if ct.globalCount >= maxGlobalConns {
|
||||
return false
|
||||
}
|
||||
if ct.perIP[ip] >= maxConnsPerIP {
|
||||
return false
|
||||
}
|
||||
|
||||
rl, ok := ct.ipRate[ip]
|
||||
if !ok {
|
||||
rl = newRateLimiter(connRateBurst, connRateSustained)
|
||||
ct.ipRate[ip] = rl
|
||||
}
|
||||
// Unlock ct.mu before calling rl.allow() would be cleaner,
|
||||
// but since rl has its own mutex this is safe (no deadlock).
|
||||
if !rl.allow() {
|
||||
return false
|
||||
}
|
||||
|
||||
ct.perIP[ip]++
|
||||
ct.globalCount++
|
||||
return true
|
||||
}
|
||||
|
||||
func (ct *connTracker) disconnect(ip string) {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
if ct.perIP[ip] > 0 {
|
||||
ct.perIP[ip]--
|
||||
ct.globalCount--
|
||||
}
|
||||
if ct.perIP[ip] == 0 {
|
||||
delete(ct.perIP, ip)
|
||||
}
|
||||
}
|
||||
|
||||
func (ct *connTracker) tryCreateRoom(ip string) bool {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
if ct.roomsPerIP[ip] >= maxRoomsPerIP {
|
||||
return false
|
||||
}
|
||||
ct.roomsPerIP[ip]++
|
||||
return true
|
||||
}
|
||||
|
||||
func (ct *connTracker) releaseRoom(ip string) {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
if ct.roomsPerIP[ip] > 0 {
|
||||
ct.roomsPerIP[ip]--
|
||||
}
|
||||
if ct.roomsPerIP[ip] == 0 {
|
||||
delete(ct.roomsPerIP, ip)
|
||||
}
|
||||
}
|
||||
|
||||
func (ct *connTracker) cleanup() {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
for ip := range ct.ipRate {
|
||||
if ct.perIP[ip] == 0 {
|
||||
delete(ct.ipRate, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Messages ---
|
||||
|
||||
type clientMsg struct {
|
||||
@@ -222,9 +84,10 @@ type serverMsg struct {
|
||||
// --- Client (serializes writes to a single goroutine) ---
|
||||
|
||||
type Client struct {
|
||||
conn *websocket.Conn
|
||||
send chan []byte
|
||||
done chan struct{}
|
||||
conn *websocket.Conn
|
||||
send chan []byte
|
||||
done chan struct{}
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
func newClient(conn *websocket.Conn) *Client {
|
||||
@@ -270,7 +133,10 @@ func (c *Client) sendJSON(msg serverMsg) {
|
||||
}
|
||||
|
||||
func (c *Client) close() {
|
||||
close(c.done)
|
||||
c.closeOnce.Do(func() {
|
||||
close(c.done)
|
||||
_ = c.conn.Close()
|
||||
})
|
||||
}
|
||||
|
||||
// --- Room ---
|
||||
@@ -312,15 +178,16 @@ func (r *Room) broadcastExcept(senderID string, msg serverMsg) {
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
// Copy peers under lock, then send without holding it
|
||||
r.mu.RLock()
|
||||
// Copy peers and record activity under lock, then send without holding it.
|
||||
r.mu.Lock()
|
||||
targets := make([]*Client, 0, len(r.Peers))
|
||||
r.LastActivityAt = time.Now()
|
||||
for id, client := range r.Peers {
|
||||
if id != senderID {
|
||||
targets = append(targets, client)
|
||||
}
|
||||
}
|
||||
r.mu.RUnlock()
|
||||
r.mu.Unlock()
|
||||
|
||||
for _, client := range targets {
|
||||
client.trySend(data)
|
||||
@@ -332,9 +199,12 @@ func (r *Room) sendTo(targetID string, msg serverMsg) bool {
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
r.mu.RLock()
|
||||
r.mu.Lock()
|
||||
client, ok := r.Peers[targetID]
|
||||
r.mu.RUnlock()
|
||||
if ok {
|
||||
r.LastActivityAt = time.Now()
|
||||
}
|
||||
r.mu.Unlock()
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
@@ -346,14 +216,18 @@ func (r *Room) sendTo(targetID string, msg serverMsg) bool {
|
||||
|
||||
type logEntry struct {
|
||||
Size int
|
||||
CreatedAt time.Time
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
var errLogStoreFull = errors.New("log store full")
|
||||
|
||||
type logStore struct {
|
||||
entries map[string]logEntry
|
||||
rateLimit map[string]time.Time // IP -> last upload time
|
||||
dir string
|
||||
mu sync.RWMutex
|
||||
entries map[string]logEntry
|
||||
rateLimit map[string]time.Time // IP -> last upload time
|
||||
dir string
|
||||
generateID func() string
|
||||
mu sync.RWMutex
|
||||
}
|
||||
|
||||
func newLogStore(dir string) *logStore {
|
||||
@@ -361,15 +235,12 @@ func newLogStore(dir string) *logStore {
|
||||
log.Fatalf("failed to create log dir %s: %v", dir, err)
|
||||
}
|
||||
ls := &logStore{
|
||||
entries: make(map[string]logEntry),
|
||||
rateLimit: make(map[string]time.Time),
|
||||
dir: dir,
|
||||
}
|
||||
// Clean orphaned files from prior runs
|
||||
files, _ := os.ReadDir(dir)
|
||||
for _, f := range files {
|
||||
os.Remove(filepath.Join(dir, f.Name()))
|
||||
entries: make(map[string]logEntry),
|
||||
rateLimit: make(map[string]time.Time),
|
||||
dir: dir,
|
||||
generateID: generateLogID,
|
||||
}
|
||||
ls.loadExisting(time.Now())
|
||||
return ls
|
||||
}
|
||||
|
||||
@@ -392,21 +263,154 @@ func generateLogID() string {
|
||||
return generateID(logIDLength)
|
||||
}
|
||||
|
||||
func logIDFromFilename(filename string) (string, bool) {
|
||||
if filepath.Ext(filename) != ".log" {
|
||||
return "", false
|
||||
}
|
||||
id := strings.TrimSuffix(filename, ".log")
|
||||
return id, validID(id, logIDLength)
|
||||
}
|
||||
|
||||
func (ls *logStore) loadExisting(now time.Time) {
|
||||
ls.mu.Lock()
|
||||
defer ls.mu.Unlock()
|
||||
|
||||
files, err := os.ReadDir(ls.dir)
|
||||
if err != nil {
|
||||
log.Printf("logs: failed to read dir %s: %v", ls.dir, err)
|
||||
return
|
||||
}
|
||||
for _, file := range files {
|
||||
filename := file.Name()
|
||||
path := filepath.Join(ls.dir, filename)
|
||||
if file.IsDir() || strings.HasSuffix(filename, ".tmp") {
|
||||
os.RemoveAll(path)
|
||||
continue
|
||||
}
|
||||
id, ok := logIDFromFilename(filename)
|
||||
if !ok {
|
||||
os.Remove(path)
|
||||
continue
|
||||
}
|
||||
info, err := file.Info()
|
||||
if err != nil || info.Size() <= 0 || info.Size() > maxLogSize {
|
||||
os.Remove(path)
|
||||
continue
|
||||
}
|
||||
createdAt := info.ModTime()
|
||||
expiresAt := createdAt.Add(logMaxAge)
|
||||
if !now.Before(expiresAt) {
|
||||
os.Remove(path)
|
||||
continue
|
||||
}
|
||||
ls.entries[id] = logEntry{
|
||||
Size: int(info.Size()),
|
||||
CreatedAt: createdAt,
|
||||
ExpiresAt: expiresAt,
|
||||
}
|
||||
}
|
||||
ls.evictOldestLocked(maxLogEntries)
|
||||
}
|
||||
|
||||
func (ls *logStore) store(data []byte, now time.Time) (string, logEntry, error) {
|
||||
if len(data) == 0 {
|
||||
return "", logEntry{}, errors.New("empty log")
|
||||
}
|
||||
if len(data) > maxLogSize {
|
||||
return "", logEntry{}, errors.New("log too large")
|
||||
}
|
||||
|
||||
ls.mu.Lock()
|
||||
defer ls.mu.Unlock()
|
||||
ls.cleanupExpiredLocked(now)
|
||||
if len(ls.entries) >= maxLogEntries {
|
||||
return "", logEntry{}, errLogStoreFull
|
||||
}
|
||||
|
||||
id := ls.generateID()
|
||||
for {
|
||||
if _, exists := ls.entries[id]; !exists {
|
||||
if _, err := os.Stat(ls.filePath(id)); errors.Is(err, fs.ErrNotExist) {
|
||||
break
|
||||
}
|
||||
}
|
||||
id = ls.generateID()
|
||||
}
|
||||
|
||||
path := ls.filePath(id)
|
||||
tmpPath := path + ".tmp"
|
||||
if err := os.WriteFile(tmpPath, data, 0644); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return "", logEntry{}, err
|
||||
}
|
||||
if err := os.Rename(tmpPath, path); err != nil {
|
||||
os.Remove(tmpPath)
|
||||
return "", logEntry{}, err
|
||||
}
|
||||
_ = os.Chtimes(path, now, now)
|
||||
|
||||
entry := logEntry{
|
||||
Size: len(data),
|
||||
CreatedAt: now,
|
||||
ExpiresAt: now.Add(logMaxAge),
|
||||
}
|
||||
ls.entries[id] = entry
|
||||
return id, entry, nil
|
||||
}
|
||||
|
||||
func (ls *logStore) lookup(id string, now time.Time) (logEntry, bool) {
|
||||
if !validID(id, logIDLength) {
|
||||
return logEntry{}, false
|
||||
}
|
||||
ls.mu.Lock()
|
||||
defer ls.mu.Unlock()
|
||||
entry, ok := ls.entries[id]
|
||||
if !ok {
|
||||
return logEntry{}, false
|
||||
}
|
||||
if !now.Before(entry.ExpiresAt) {
|
||||
ls.deleteEntryLocked(id)
|
||||
return logEntry{}, false
|
||||
}
|
||||
return entry, true
|
||||
}
|
||||
|
||||
func (ls *logStore) cleanupExpiredLocked(now time.Time) {
|
||||
for id, entry := range ls.entries {
|
||||
if !now.Before(entry.ExpiresAt) {
|
||||
ls.deleteEntryLocked(id)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (ls *logStore) evictOldestLocked(limit int) {
|
||||
for len(ls.entries) > limit {
|
||||
var oldestID string
|
||||
var oldest logEntry
|
||||
for id, entry := range ls.entries {
|
||||
if oldestID == "" || entry.CreatedAt.Before(oldest.CreatedAt) {
|
||||
oldestID = id
|
||||
oldest = entry
|
||||
}
|
||||
}
|
||||
if oldestID == "" {
|
||||
return
|
||||
}
|
||||
ls.deleteEntryLocked(oldestID)
|
||||
}
|
||||
}
|
||||
|
||||
func (ls *logStore) deleteEntryLocked(id string) {
|
||||
os.Remove(ls.filePath(id))
|
||||
delete(ls.entries, id)
|
||||
}
|
||||
|
||||
func (ls *logStore) cleanup() {
|
||||
ls.mu.Lock()
|
||||
defer ls.mu.Unlock()
|
||||
now := time.Now()
|
||||
for id, entry := range ls.entries {
|
||||
if now.After(entry.ExpiresAt) {
|
||||
os.Remove(ls.filePath(id))
|
||||
delete(ls.entries, id)
|
||||
}
|
||||
}
|
||||
for ip, lastTime := range ls.rateLimit {
|
||||
if now.Sub(lastTime) > logRateInterval {
|
||||
delete(ls.rateLimit, ip)
|
||||
}
|
||||
}
|
||||
ls.cleanupExpiredLocked(now)
|
||||
cleanupRateWindows(ls.rateLimit, now, logRateInterval)
|
||||
}
|
||||
|
||||
// --- Poster store ---
|
||||
@@ -903,7 +907,7 @@ func (s *Server) loadSnapshot(path string) error {
|
||||
loaded, skipped := 0, 0
|
||||
s.mu.Lock()
|
||||
for _, r := range snap.Rooms {
|
||||
if r.SessionID == "" || r.HostPeerID == "" {
|
||||
if !validRelayID(r.SessionID, maxSessionIDLength) || !validRelayID(r.HostPeerID, maxPeerIDLength) {
|
||||
skipped++
|
||||
continue
|
||||
}
|
||||
@@ -940,33 +944,45 @@ func (s *Server) cleanupLoop() {
|
||||
func (s *Server) runCleanupStep(now time.Time) {
|
||||
s.mu.Lock()
|
||||
changed := false
|
||||
var expiredClients []*Client
|
||||
for id, room := range s.rooms {
|
||||
room.mu.RLock()
|
||||
empty := len(room.Peers) == 0
|
||||
age := now.Sub(room.CreatedAt)
|
||||
idle := now.Sub(room.LastActivityAt)
|
||||
expired := age > roomMaxAge
|
||||
if expired && !empty {
|
||||
for _, client := range room.Peers {
|
||||
expiredClients = append(expiredClients, client)
|
||||
}
|
||||
}
|
||||
room.mu.RUnlock()
|
||||
|
||||
if (empty && idle > emptyRoomMaxAge) || age > roomMaxAge {
|
||||
if (empty && idle > emptyRoomMaxAge) || expired {
|
||||
log.Printf("cleanup: removing room %s (empty=%v, idle=%v, age=%v)", id, empty, idle, age)
|
||||
delete(s.rooms, id)
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
roomCount := len(s.rooms)
|
||||
s.mu.Unlock()
|
||||
|
||||
for _, client := range expiredClients {
|
||||
client.close()
|
||||
}
|
||||
if changed {
|
||||
s.snap.schedule()
|
||||
}
|
||||
s.logs.cleanup()
|
||||
s.posters.cleanup(now)
|
||||
s.conns.cleanup()
|
||||
s.conns.cleanup(now)
|
||||
if s.oauth != nil {
|
||||
s.oauth.cleanup()
|
||||
}
|
||||
|
||||
s.conns.mu.Lock()
|
||||
log.Printf("stats: conns=%d ips=%d rooms=%d",
|
||||
s.conns.globalCount, len(s.conns.perIP), len(s.rooms))
|
||||
s.conns.globalCount, len(s.conns.perIP), roomCount)
|
||||
s.conns.mu.Unlock()
|
||||
}
|
||||
|
||||
@@ -1021,29 +1037,18 @@ func (s *Server) handlePostLogs(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
s.logs.mu.Lock()
|
||||
if len(s.logs.entries) >= maxLogEntries {
|
||||
s.logs.mu.Unlock()
|
||||
http.Error(w, "Log store full", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
s.logs.mu.Unlock()
|
||||
|
||||
id := generateLogID()
|
||||
if err := os.WriteFile(s.logs.filePath(id), body, 0644); err != nil {
|
||||
log.Printf("logs: failed to write %s: %v", id, err)
|
||||
id, entry, err := s.logs.store(body, time.Now())
|
||||
if err != nil {
|
||||
if errors.Is(err, errLogStoreFull) {
|
||||
http.Error(w, "Log store full", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
log.Printf("logs: failed to store from %s: %v", ip, err)
|
||||
http.Error(w, "Failed to store log", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
s.logs.mu.Lock()
|
||||
s.logs.entries[id] = logEntry{
|
||||
Size: len(body),
|
||||
ExpiresAt: time.Now().Add(logMaxAge),
|
||||
}
|
||||
s.logs.mu.Unlock()
|
||||
|
||||
log.Printf("logs: stored %s (%d bytes) from %s", id, len(body), ip)
|
||||
log.Printf("logs: stored %s (%d bytes) from %s", id, entry.Size, ip)
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]string{"id": id})
|
||||
@@ -1061,11 +1066,8 @@ func (s *Server) handleGetLogs(w http.ResponseWriter, r *http.Request) {
|
||||
return
|
||||
}
|
||||
|
||||
s.logs.mu.RLock()
|
||||
entry, ok := s.logs.entries[id]
|
||||
s.logs.mu.RUnlock()
|
||||
|
||||
if !ok || time.Now().After(entry.ExpiresAt) {
|
||||
entry, ok := s.logs.lookup(id, time.Now())
|
||||
if !ok {
|
||||
http.Error(w, "Not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
@@ -1185,6 +1187,17 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) {
|
||||
var currentRoom *Room
|
||||
var currentPeerID string
|
||||
var isHost bool
|
||||
rejectRoomTransition := func() bool {
|
||||
if currentRoom == nil {
|
||||
return false
|
||||
}
|
||||
client.sendJSON(serverMsg{
|
||||
Type: relayTypeError,
|
||||
Code: relayErrorAlreadyInRoom,
|
||||
Message: "Leave the current room before creating or joining another",
|
||||
})
|
||||
return true
|
||||
}
|
||||
|
||||
// Cleanup on disconnect — only if our Client is still the one in the room.
|
||||
// A reconnecting peer reuses the same peerId, so the map entry may have
|
||||
@@ -1200,7 +1213,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) {
|
||||
currentRoom.mu.Unlock()
|
||||
if !stale {
|
||||
currentRoom.broadcastExcept(currentPeerID, serverMsg{
|
||||
Type: "peerLeft",
|
||||
Type: relayTypePeerLeft,
|
||||
PeerID: currentPeerID,
|
||||
})
|
||||
s.snap.schedule()
|
||||
@@ -1222,24 +1235,27 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
|
||||
if !rl.allow() {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "rate_limited", Message: "Too many messages"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRateLimited, Message: "Too many messages"})
|
||||
continue
|
||||
}
|
||||
|
||||
var msg clientMsg
|
||||
if err := json.Unmarshal(raw, &msg); err != nil {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "invalid_message", Message: "Invalid JSON"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid JSON"})
|
||||
continue
|
||||
}
|
||||
|
||||
switch msg.Type {
|
||||
case "create":
|
||||
if msg.SessionID == "" || msg.PeerID == "" {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "invalid_message", Message: "sessionId and peerId required"})
|
||||
case relayTypeCreate:
|
||||
if !validRelayID(msg.SessionID, maxSessionIDLength) || !validRelayID(msg.PeerID, maxPeerIDLength) {
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid sessionId or peerId"})
|
||||
continue
|
||||
}
|
||||
if rejectRoomTransition() {
|
||||
continue
|
||||
}
|
||||
if !s.conns.tryCreateRoom(ip) {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "rate_limited", Message: "Too many rooms created"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRateLimited, Message: "Too many rooms created"})
|
||||
continue
|
||||
}
|
||||
s.mu.Lock()
|
||||
@@ -1250,7 +1266,7 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) {
|
||||
if !empty {
|
||||
s.mu.Unlock()
|
||||
s.conns.releaseRoom(ip)
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "room_exists", Message: "Room already exists"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomExists, Message: "Room already exists"})
|
||||
continue
|
||||
}
|
||||
// Empty stale room — reclaim the ID
|
||||
@@ -1270,25 +1286,28 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) {
|
||||
currentPeerID = msg.PeerID
|
||||
isHost = true
|
||||
log.Printf("room %s created by %s", msg.SessionID, msg.PeerID)
|
||||
client.sendJSON(serverMsg{Type: "created", SessionID: msg.SessionID})
|
||||
client.sendJSON(serverMsg{Type: relayTypeCreated, SessionID: msg.SessionID})
|
||||
s.snap.schedule()
|
||||
|
||||
case "join":
|
||||
if msg.SessionID == "" || msg.PeerID == "" {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "invalid_message", Message: "sessionId and peerId required"})
|
||||
case relayTypeJoin:
|
||||
if !validRelayID(msg.SessionID, maxSessionIDLength) || !validRelayID(msg.PeerID, maxPeerIDLength) {
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid sessionId or peerId"})
|
||||
continue
|
||||
}
|
||||
if rejectRoomTransition() {
|
||||
continue
|
||||
}
|
||||
s.mu.RLock()
|
||||
room, exists := s.rooms[msg.SessionID]
|
||||
s.mu.RUnlock()
|
||||
if !exists {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "room_not_found", Message: "Room does not exist"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomNotFound, Message: "Room does not exist"})
|
||||
continue
|
||||
}
|
||||
room.mu.Lock()
|
||||
if len(room.Peers) >= maxRoomSize {
|
||||
room.mu.Unlock()
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "room_full", Message: "Room is full"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomFull, Message: "Room is full"})
|
||||
continue
|
||||
}
|
||||
room.Peers[msg.PeerID] = client
|
||||
@@ -1306,43 +1325,43 @@ func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) {
|
||||
existingPeers = append(existingPeers, p)
|
||||
}
|
||||
}
|
||||
client.sendJSON(serverMsg{Type: "joined", SessionID: msg.SessionID, Peers: existingPeers})
|
||||
room.broadcastExcept(msg.PeerID, serverMsg{Type: "peerJoined", PeerID: msg.PeerID})
|
||||
client.sendJSON(serverMsg{Type: relayTypeJoined, SessionID: msg.SessionID, Peers: existingPeers})
|
||||
room.broadcastExcept(msg.PeerID, serverMsg{Type: relayTypePeerJoined, PeerID: msg.PeerID})
|
||||
s.snap.schedule()
|
||||
|
||||
case "broadcast":
|
||||
case relayTypeBroadcast:
|
||||
if currentRoom == nil {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "not_in_room", Message: "Not in a room"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"})
|
||||
continue
|
||||
}
|
||||
currentRoom.broadcastExcept(currentPeerID, serverMsg{
|
||||
Type: "message",
|
||||
Type: relayTypeMessage,
|
||||
From: currentPeerID,
|
||||
Payload: msg.Payload,
|
||||
})
|
||||
|
||||
case "sendTo":
|
||||
case relayTypeSendTo:
|
||||
if currentRoom == nil {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "not_in_room", Message: "Not in a room"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"})
|
||||
continue
|
||||
}
|
||||
if msg.To == "" {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "invalid_message", Message: "to field required"})
|
||||
if !validRelayID(msg.To, maxPeerIDLength) {
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid to field"})
|
||||
continue
|
||||
}
|
||||
if !currentRoom.sendTo(msg.To, serverMsg{
|
||||
Type: "message",
|
||||
Type: relayTypeMessage,
|
||||
From: currentPeerID,
|
||||
Payload: msg.Payload,
|
||||
}) {
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "not_in_room", Message: "Target peer not found"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Target peer not found"})
|
||||
}
|
||||
|
||||
case "ping":
|
||||
client.sendJSON(serverMsg{Type: "pong"})
|
||||
case relayTypePing:
|
||||
client.sendJSON(serverMsg{Type: relayTypePong})
|
||||
|
||||
default:
|
||||
client.sendJSON(serverMsg{Type: "error", Code: "invalid_message", Message: "Unknown message type"})
|
||||
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unknown message type"})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+296
-13
@@ -392,6 +392,21 @@ func (h *relayHarness) waitRoomPeers(t *testing.T, sessionID string, want int) {
|
||||
t.Fatalf("room %s never reached %d peers within 2s", sessionID, want)
|
||||
}
|
||||
|
||||
func (h *relayHarness) waitIPConnections(t *testing.T, ip string, want int) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
h.srv.conns.mu.Lock()
|
||||
got := h.srv.conns.perIP[ip]
|
||||
h.srv.conns.mu.Unlock()
|
||||
if got == want {
|
||||
return
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
t.Fatalf("IP %s never reached %d connections within 2s", ip, want)
|
||||
}
|
||||
|
||||
type testConn struct {
|
||||
t *testing.T
|
||||
conn *websocket.Conn
|
||||
@@ -521,6 +536,40 @@ func TestRateLimiterAllowRace(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRateLimiterReclaimableOnlyAfterFullRefill(t *testing.T) {
|
||||
now := time.Now()
|
||||
limiter := &rateLimiter{
|
||||
tokens: 0,
|
||||
maxTokens: 5,
|
||||
refillRate: 1,
|
||||
lastTime: now,
|
||||
}
|
||||
|
||||
if limiter.reclaimable(now.Add(4 * time.Second)) {
|
||||
t.Fatal("partially refilled limiter must retain its effective state")
|
||||
}
|
||||
if !limiter.reclaimable(now.Add(5 * time.Second)) {
|
||||
t.Fatal("fully refilled limiter should be reclaimable")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupRateWindowsUsesWindowBoundary(t *testing.T) {
|
||||
now := time.Now()
|
||||
windows := map[string]time.Time{
|
||||
"active": now.Add(-logRateInterval + time.Nanosecond),
|
||||
"expired": now.Add(-logRateInterval),
|
||||
}
|
||||
|
||||
cleanupRateWindows(windows, now, logRateInterval)
|
||||
|
||||
if _, ok := windows["active"]; !ok {
|
||||
t.Fatal("active fixed-window limiter was removed early")
|
||||
}
|
||||
if _, ok := windows["expired"]; ok {
|
||||
t.Fatal("expired fixed-window limiter was retained")
|
||||
}
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// connTracker unit tests
|
||||
// ======================================================================
|
||||
@@ -588,25 +637,34 @@ func TestConnTrackerRoomQuota(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnTrackerCleanupPrunesStaleRateLimiters(t *testing.T) {
|
||||
func TestConnTrackerCleanupPreservesEffectiveRateLimits(t *testing.T) {
|
||||
ct := newConnTracker()
|
||||
for i := 0; i < 50; i++ {
|
||||
ip := fmt.Sprintf("10.0.1.%d", i)
|
||||
ct.tryConnect(ip)
|
||||
ip := "10.0.1.1"
|
||||
for i := range connRateBurst {
|
||||
if !ct.tryConnect(ip) {
|
||||
t.Fatalf("tryConnect %d: expected true", i)
|
||||
}
|
||||
}
|
||||
for range connRateBurst {
|
||||
ct.disconnect(ip)
|
||||
}
|
||||
ct.mu.Lock()
|
||||
sizeBefore := len(ct.ipRate)
|
||||
ct.mu.Unlock()
|
||||
if sizeBefore == 0 {
|
||||
t.Fatal("expected some rate limiter entries before cleanup")
|
||||
|
||||
ct.cleanup(time.Now())
|
||||
if ct.tryConnect(ip) {
|
||||
t.Fatal("cleanup reset a connection rate limit that was still effective")
|
||||
}
|
||||
ct.cleanup()
|
||||
|
||||
ct.cleanup(time.Now().Add(10 * time.Second))
|
||||
if !ct.tryConnect(ip) {
|
||||
t.Fatal("fully refilled limiter should be reclaimable")
|
||||
}
|
||||
|
||||
ct.cleanup(time.Now().Add(10 * time.Second))
|
||||
ct.mu.Lock()
|
||||
sizeAfter := len(ct.ipRate)
|
||||
_, retainedWhileConnected := ct.ipRate[ip]
|
||||
ct.mu.Unlock()
|
||||
if sizeAfter != 0 {
|
||||
t.Errorf("cleanup should prune all stale rate limiters, got %d", sizeAfter)
|
||||
if !retainedWhileConnected {
|
||||
t.Fatal("cleanup removed a limiter with an active connection")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -771,6 +829,39 @@ func TestCreateHitsRoomsPerIPLimit(t *testing.T) {
|
||||
c.expectError("rate_limited")
|
||||
}
|
||||
|
||||
func TestConnectionCannotRetainMultipleRoomMemberships(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
ip := "1.1.1.8"
|
||||
client := h.dial(t, ip)
|
||||
client.send(clientMsg{Type: "create", SessionID: "PRIMARY", PeerID: "host"})
|
||||
client.expect("created")
|
||||
|
||||
otherHost := h.dial(t, "1.1.1.9")
|
||||
otherHost.send(clientMsg{Type: "create", SessionID: "OTHER", PeerID: "other-host"})
|
||||
otherHost.expect("created")
|
||||
|
||||
client.send(clientMsg{Type: "join", SessionID: "OTHER", PeerID: "ghost"})
|
||||
client.expectError("already_in_room")
|
||||
client.send(clientMsg{Type: "create", SessionID: "EXTRA", PeerID: "extra-host"})
|
||||
client.expectError("already_in_room")
|
||||
|
||||
h.waitRoomPeers(t, "PRIMARY", 1)
|
||||
h.waitRoomPeers(t, "OTHER", 1)
|
||||
h.srv.mu.RLock()
|
||||
_, extraExists := h.srv.rooms["EXTRA"]
|
||||
h.srv.mu.RUnlock()
|
||||
if extraExists {
|
||||
t.Fatal("rejected create retained an extra room")
|
||||
}
|
||||
|
||||
h.srv.conns.mu.Lock()
|
||||
roomsForIP := h.srv.conns.roomsPerIP[ip]
|
||||
h.srv.conns.mu.Unlock()
|
||||
if roomsForIP != 1 {
|
||||
t.Fatalf("roomsPerIP[%q]=%d, want 1", ip, roomsForIP)
|
||||
}
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// handleWS — join case
|
||||
// ======================================================================
|
||||
@@ -805,6 +896,22 @@ func TestJoinMissingFieldsRejected(t *testing.T) {
|
||||
c.expectError("invalid_message")
|
||||
}
|
||||
|
||||
func TestRelayIdentifiersRejectUnsafeOrOversizedValues(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
c := h.dial(t, "2.0.0.30")
|
||||
|
||||
invalid := []string{"has space", "has/slash", strings.Repeat("x", maxSessionIDLength+1)}
|
||||
for _, sessionID := range invalid {
|
||||
c.send(clientMsg{Type: "create", SessionID: sessionID, PeerID: "H"})
|
||||
c.expectError("invalid_message")
|
||||
}
|
||||
|
||||
c.send(clientMsg{Type: "create", SessionID: "SAFE_ID-1", PeerID: "H"})
|
||||
c.expect("created")
|
||||
c.send(clientMsg{Type: "sendTo", To: "bad target", Payload: json.RawMessage(`{}`)})
|
||||
c.expectError("invalid_message")
|
||||
}
|
||||
|
||||
func TestJoinUnknownRoomFails(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
c := h.dial(t, "2.0.0.4")
|
||||
@@ -907,6 +1014,47 @@ func TestSendToDeliversToTargetOnly(t *testing.T) {
|
||||
g2.recvNothing(200 * time.Millisecond)
|
||||
}
|
||||
|
||||
func TestRelayMessagesRefreshRoomActivity(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
host := h.dial(t, "4.0.0.7")
|
||||
host.send(clientMsg{Type: "create", SessionID: "ACTIVE", PeerID: "H"})
|
||||
host.expect("created")
|
||||
|
||||
guest := h.dial(t, "4.0.0.8")
|
||||
guest.send(clientMsg{Type: "join", SessionID: "ACTIVE", PeerID: "G"})
|
||||
guest.expect("joined")
|
||||
host.expect("peerJoined")
|
||||
|
||||
h.srv.mu.RLock()
|
||||
room := h.srv.rooms["ACTIVE"]
|
||||
h.srv.mu.RUnlock()
|
||||
old := time.Now().Add(-time.Hour)
|
||||
|
||||
room.mu.Lock()
|
||||
room.LastActivityAt = old
|
||||
room.mu.Unlock()
|
||||
host.send(clientMsg{Type: "broadcast", Payload: json.RawMessage(`{"broadcast":true}`)})
|
||||
guest.expect("message")
|
||||
room.mu.RLock()
|
||||
broadcastActivity := room.LastActivityAt
|
||||
room.mu.RUnlock()
|
||||
if !broadcastActivity.After(old) {
|
||||
t.Fatalf("broadcast activity=%v, want after %v", broadcastActivity, old)
|
||||
}
|
||||
|
||||
room.mu.Lock()
|
||||
room.LastActivityAt = old
|
||||
room.mu.Unlock()
|
||||
host.send(clientMsg{Type: "sendTo", To: "G", Payload: json.RawMessage(`{"direct":true}`)})
|
||||
guest.expect("message")
|
||||
room.mu.RLock()
|
||||
directActivity := room.LastActivityAt
|
||||
room.mu.RUnlock()
|
||||
if !directActivity.After(old) {
|
||||
t.Fatalf("sendTo activity=%v, want after %v", directActivity, old)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendToUnknownTargetRejected(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
host := h.dial(t, "4.0.0.4")
|
||||
@@ -1027,6 +1175,103 @@ func TestStalePeerSkipsCleanupBroadcast(t *testing.T) {
|
||||
host.recvNothing(300 * time.Millisecond)
|
||||
}
|
||||
|
||||
func TestHostReconnectReplacesStaleConnectionWithoutLeaving(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
oldHostIP := "6.1.1.1"
|
||||
oldHost := h.dial(t, oldHostIP)
|
||||
oldHost.send(clientMsg{Type: "create", SessionID: "REJOIN", PeerID: "H"})
|
||||
oldHost.expect("created")
|
||||
|
||||
guest := h.dial(t, "6.1.1.2")
|
||||
guest.send(clientMsg{Type: "join", SessionID: "REJOIN", PeerID: "G"})
|
||||
guest.expect("joined")
|
||||
oldHost.expect("peerJoined")
|
||||
|
||||
newHost := h.dial(t, "6.1.1.3")
|
||||
newHost.send(clientMsg{Type: "join", SessionID: "REJOIN", PeerID: "H"})
|
||||
joined := newHost.expect("joined")
|
||||
if len(joined.Peers) != 1 || joined.Peers[0] != "G" {
|
||||
t.Fatalf("reconnected host peers=%v, want [G]", joined.Peers)
|
||||
}
|
||||
guest.expect("peerJoined")
|
||||
|
||||
oldHost.conn.Close()
|
||||
h.waitIPConnections(t, oldHostIP, 0)
|
||||
|
||||
newHost.send(clientMsg{Type: "broadcast", Payload: json.RawMessage(`{"state":"ready"}`)})
|
||||
message := guest.expect("message")
|
||||
if message.From != "H" {
|
||||
t.Fatalf("message sender=%q, want H", message.From)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmptyRoomSupportsJoinThenExpiresForReconnectFallback(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
host := h.dial(t, "6.1.2.1")
|
||||
host.send(clientMsg{Type: "create", SessionID: "EMPTY", PeerID: "H"})
|
||||
host.expect("created")
|
||||
host.conn.Close()
|
||||
h.waitRoomPeers(t, "EMPTY", 0)
|
||||
|
||||
reconnected := h.dial(t, "6.1.2.2")
|
||||
reconnected.send(clientMsg{Type: "join", SessionID: "EMPTY", PeerID: "H"})
|
||||
joined := reconnected.expect("joined")
|
||||
if len(joined.Peers) != 0 {
|
||||
t.Fatalf("empty-room reconnect peers=%v, want none", joined.Peers)
|
||||
}
|
||||
reconnected.conn.Close()
|
||||
h.waitRoomPeers(t, "EMPTY", 0)
|
||||
|
||||
now := time.Now()
|
||||
h.srv.mu.RLock()
|
||||
room := h.srv.rooms["EMPTY"]
|
||||
h.srv.mu.RUnlock()
|
||||
room.mu.Lock()
|
||||
room.LastActivityAt = now.Add(-emptyRoomMaxAge - time.Second)
|
||||
room.mu.Unlock()
|
||||
h.srv.runCleanupStep(now)
|
||||
|
||||
fallback := h.dial(t, "6.1.2.3")
|
||||
fallback.send(clientMsg{Type: "join", SessionID: "EMPTY", PeerID: "H"})
|
||||
fallback.expectError("room_not_found")
|
||||
}
|
||||
|
||||
func TestCleanupDisconnectsPeersBeforeRemovingExpiredOccupiedRoom(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
host := h.dial(t, "6.2.0.1")
|
||||
host.send(clientMsg{Type: "create", SessionID: "EXPIRED", PeerID: "H"})
|
||||
host.expect("created")
|
||||
|
||||
guest := h.dial(t, "6.2.0.2")
|
||||
guest.send(clientMsg{Type: "join", SessionID: "EXPIRED", PeerID: "G"})
|
||||
guest.expect("joined")
|
||||
host.expect("peerJoined")
|
||||
|
||||
now := time.Now()
|
||||
h.srv.mu.RLock()
|
||||
room := h.srv.rooms["EXPIRED"]
|
||||
h.srv.mu.RUnlock()
|
||||
room.mu.Lock()
|
||||
room.CreatedAt = now.Add(-roomMaxAge - time.Second)
|
||||
room.mu.Unlock()
|
||||
|
||||
h.srv.runCleanupStep(now)
|
||||
|
||||
h.srv.mu.RLock()
|
||||
_, exists := h.srv.rooms["EXPIRED"]
|
||||
h.srv.mu.RUnlock()
|
||||
if exists {
|
||||
t.Fatal("expired room still exists after cleanup")
|
||||
}
|
||||
|
||||
for name, connection := range map[string]*testConn{"host": host, "guest": guest} {
|
||||
connection.conn.SetReadDeadline(time.Now().Add(2 * time.Second))
|
||||
if _, _, err := connection.conn.ReadMessage(); err == nil {
|
||||
t.Errorf("%s remained connected after occupied room removal", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ======================================================================
|
||||
// Logs endpoints
|
||||
// ======================================================================
|
||||
@@ -1132,6 +1377,44 @@ func TestLogsRoundTrip(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogStorePersistsAcrossRestartAndAvoidsIDCollisions(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
now := time.Now().Add(-time.Second)
|
||||
first := newLogStore(dir)
|
||||
first.generateID = func() string { return "aaaaa" }
|
||||
firstID, _, err := first.store([]byte("original"), now)
|
||||
if err != nil {
|
||||
t.Fatalf("store original: %v", err)
|
||||
}
|
||||
|
||||
restarted := newLogStore(dir)
|
||||
if _, ok := restarted.lookup(firstID, time.Now()); !ok {
|
||||
t.Fatal("stored log was not restored after restart")
|
||||
}
|
||||
|
||||
ids := []string{firstID, "bbbbb"}
|
||||
restarted.generateID = func() string {
|
||||
id := ids[0]
|
||||
ids = ids[1:]
|
||||
return id
|
||||
}
|
||||
secondID, _, err := restarted.store([]byte("second"), time.Now())
|
||||
if err != nil {
|
||||
t.Fatalf("store after restart: %v", err)
|
||||
}
|
||||
if secondID != "bbbbb" {
|
||||
t.Fatalf("collision generated id %q, want bbbbb", secondID)
|
||||
}
|
||||
|
||||
original, err := os.ReadFile(restarted.filePath(firstID))
|
||||
if err != nil {
|
||||
t.Fatalf("read original: %v", err)
|
||||
}
|
||||
if string(original) != "original" {
|
||||
t.Fatalf("colliding store overwrote original: %q", original)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogsUploadRateLimitedPerIP(t *testing.T) {
|
||||
h := newRelayHarness(t)
|
||||
r1 := postLog(t, h.baseURL, "7.0.0.2", []byte("first"))
|
||||
|
||||
+8
-12
@@ -25,14 +25,14 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
oauthSessionTTL = 10 * time.Minute
|
||||
oauthResultWait = 50 * time.Second
|
||||
oauthMaxSessions = 5000
|
||||
oauthStartBurst = 3
|
||||
oauthSessionTTL = 10 * time.Minute
|
||||
oauthResultWait = 50 * time.Second
|
||||
oauthMaxSessions = 5000
|
||||
oauthStartBurst = 3
|
||||
oauthStartRateSustained = 1
|
||||
oauthSessionIDBytes = 18 // 144 bits → 24 base64url chars
|
||||
oauthPKCEVerifierLen = 64
|
||||
oauthUpstreamTimeout = 15 * time.Second
|
||||
oauthSessionIDBytes = 18 // 144 bits → 24 base64url chars
|
||||
oauthPKCEVerifierLen = 64
|
||||
oauthUpstreamTimeout = 15 * time.Second
|
||||
)
|
||||
|
||||
// oauthServiceConfig describes a single upstream OAuth provider. Populated from
|
||||
@@ -422,11 +422,7 @@ func (p *oauthProxy) cleanup() {
|
||||
p.mu.Unlock()
|
||||
|
||||
p.ipMu.Lock()
|
||||
for ip, rl := range p.ipRate {
|
||||
if rl.stale(now) {
|
||||
delete(p.ipRate, ip)
|
||||
}
|
||||
}
|
||||
cleanupRateLimiters(p.ipRate, now, nil)
|
||||
p.ipMu.Unlock()
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// --- Rate limiter (token bucket) ---
|
||||
|
||||
type rateLimiter struct {
|
||||
tokens float64
|
||||
maxTokens float64
|
||||
refillRate float64
|
||||
lastTime time.Time
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func newRateLimiter(burst, sustained int) *rateLimiter {
|
||||
return &rateLimiter{
|
||||
tokens: float64(burst),
|
||||
maxTokens: float64(burst),
|
||||
refillRate: float64(sustained),
|
||||
lastTime: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
func (rl *rateLimiter) allow() bool {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
elapsed := now.Sub(rl.lastTime).Seconds()
|
||||
rl.lastTime = now
|
||||
|
||||
rl.tokens += elapsed * rl.refillRate
|
||||
if rl.tokens > rl.maxTokens {
|
||||
rl.tokens = rl.maxTokens
|
||||
}
|
||||
|
||||
if rl.tokens < 1 {
|
||||
return false
|
||||
}
|
||||
rl.tokens--
|
||||
return true
|
||||
}
|
||||
|
||||
// reclaimable reports whether discarding this limiter would preserve its
|
||||
// behavior: enough idle time has passed for the bucket to be full again.
|
||||
func (rl *rateLimiter) reclaimable(now time.Time) bool {
|
||||
rl.mu.Lock()
|
||||
defer rl.mu.Unlock()
|
||||
missingTokens := rl.maxTokens - rl.tokens
|
||||
return missingTokens <= 0 || now.Sub(rl.lastTime).Seconds()*rl.refillRate >= missingTokens
|
||||
}
|
||||
|
||||
func cleanupRateLimiters(limiters map[string]*rateLimiter, now time.Time, inUse func(string) bool) {
|
||||
for ip, limiter := range limiters {
|
||||
if (inUse == nil || !inUse(ip)) && limiter.reclaimable(now) {
|
||||
delete(limiters, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func cleanupRateWindows(windows map[string]time.Time, now time.Time, duration time.Duration) {
|
||||
for ip, startedAt := range windows {
|
||||
if now.Sub(startedAt) >= duration {
|
||||
delete(windows, ip)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// --- Connection tracker (per-IP limits) ---
|
||||
|
||||
type connTracker struct {
|
||||
mu sync.Mutex
|
||||
perIP map[string]int
|
||||
ipRate map[string]*rateLimiter
|
||||
roomsPerIP map[string]int
|
||||
globalCount int
|
||||
}
|
||||
|
||||
func newConnTracker() *connTracker {
|
||||
return &connTracker{
|
||||
perIP: make(map[string]int),
|
||||
ipRate: make(map[string]*rateLimiter),
|
||||
roomsPerIP: make(map[string]int),
|
||||
}
|
||||
}
|
||||
|
||||
func (ct *connTracker) tryConnect(ip string) bool {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
if ct.globalCount >= maxGlobalConns {
|
||||
return false
|
||||
}
|
||||
if ct.perIP[ip] >= maxConnsPerIP {
|
||||
return false
|
||||
}
|
||||
|
||||
rl, ok := ct.ipRate[ip]
|
||||
if !ok {
|
||||
rl = newRateLimiter(connRateBurst, connRateSustained)
|
||||
ct.ipRate[ip] = rl
|
||||
}
|
||||
// Unlock ct.mu before calling rl.allow() would be cleaner,
|
||||
// but since rl has its own mutex this is safe (no deadlock).
|
||||
if !rl.allow() {
|
||||
return false
|
||||
}
|
||||
|
||||
ct.perIP[ip]++
|
||||
ct.globalCount++
|
||||
return true
|
||||
}
|
||||
|
||||
func (ct *connTracker) disconnect(ip string) {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
|
||||
if ct.perIP[ip] > 0 {
|
||||
ct.perIP[ip]--
|
||||
ct.globalCount--
|
||||
}
|
||||
if ct.perIP[ip] == 0 {
|
||||
delete(ct.perIP, ip)
|
||||
}
|
||||
}
|
||||
|
||||
func (ct *connTracker) tryCreateRoom(ip string) bool {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
if ct.roomsPerIP[ip] >= maxRoomsPerIP {
|
||||
return false
|
||||
}
|
||||
ct.roomsPerIP[ip]++
|
||||
return true
|
||||
}
|
||||
|
||||
func (ct *connTracker) releaseRoom(ip string) {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
if ct.roomsPerIP[ip] > 0 {
|
||||
ct.roomsPerIP[ip]--
|
||||
}
|
||||
if ct.roomsPerIP[ip] == 0 {
|
||||
delete(ct.roomsPerIP, ip)
|
||||
}
|
||||
}
|
||||
|
||||
func (ct *connTracker) cleanup(now time.Time) {
|
||||
ct.mu.Lock()
|
||||
defer ct.mu.Unlock()
|
||||
cleanupRateLimiters(ct.ipRate, now, func(ip string) bool {
|
||||
return ct.perIP[ip] > 0
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
// Code generated by scripts/generate_relay_protocol.py. DO NOT EDIT.
|
||||
|
||||
package main
|
||||
|
||||
const (
|
||||
relayTypeCreate = "create"
|
||||
relayTypeJoin = "join"
|
||||
relayTypeBroadcast = "broadcast"
|
||||
relayTypeSendTo = "sendTo"
|
||||
relayTypePing = "ping"
|
||||
relayTypeCreated = "created"
|
||||
relayTypeJoined = "joined"
|
||||
relayTypePeerJoined = "peerJoined"
|
||||
relayTypePeerLeft = "peerLeft"
|
||||
relayTypeMessage = "message"
|
||||
relayTypeError = "error"
|
||||
relayTypePong = "pong"
|
||||
relayErrorRateLimited = "rate_limited"
|
||||
relayErrorInvalidMessage = "invalid_message"
|
||||
relayErrorRoomExists = "room_exists"
|
||||
relayErrorRoomNotFound = "room_not_found"
|
||||
relayErrorRoomFull = "room_full"
|
||||
relayErrorNotInRoom = "not_in_room"
|
||||
relayErrorAlreadyInRoom = "already_in_room"
|
||||
|
||||
maxRoomSize = 8
|
||||
maxMessageSize = 65536
|
||||
maxSessionIDLength = 64
|
||||
maxPeerIDLength = 128
|
||||
)
|
||||
|
||||
func validRelayID(value string, maxLength int) bool {
|
||||
if len(value) == 0 || len(value) > maxLength {
|
||||
return false
|
||||
}
|
||||
for _, ch := range value {
|
||||
if (ch < 'a' || ch > 'z') && (ch < 'A' || ch > 'Z') &&
|
||||
(ch < '0' || ch > '9') && ch != '_' && ch != '-' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
Reference in New Issue
Block a user