A 25-character upload capability is unreadable over the phone or in a support thread, which is the only way these ids are ever exchanged. Lookups stay bounded by the per-source failed-lookup limiter and the three-day expiry, and ids minted at the longer shape are retired on the next startup because they no longer match the store's filename shape.
2417 lines
70 KiB
Go
2417 lines
70 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"crypto/subtle"
|
|
"encoding/base64"
|
|
"encoding/json"
|
|
"errors"
|
|
"flag"
|
|
"fmt"
|
|
"io"
|
|
"io/fs"
|
|
"log"
|
|
"net"
|
|
"net/http"
|
|
"os"
|
|
"os/signal"
|
|
"path/filepath"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"syscall"
|
|
"time"
|
|
|
|
"github.com/gorilla/websocket"
|
|
)
|
|
|
|
const (
|
|
rateBurst = 30
|
|
rateSustained = 10
|
|
cleanupInterval = 5 * time.Minute
|
|
emptyRoomMaxAge = 5 * time.Minute
|
|
peerReservationGrace = emptyRoomMaxAge
|
|
roomMaxAge = 24 * time.Hour
|
|
writeWait = 10 * time.Second
|
|
httpResponseWriteMargin = 10 * time.Second
|
|
httpResponseWriteTimeout = oauthResultWait + httpResponseWriteMargin
|
|
pongWait = 60 * time.Second
|
|
pingInterval = 30 * time.Second
|
|
maxLogSize = 1 * 1024 * 1024 // 1MB
|
|
logMaxAge = 3 * 24 * time.Hour
|
|
logIDLength = 5
|
|
logRateInterval = 1 * time.Minute
|
|
logLookupRateBurst = 10
|
|
logLookupRateSustained = 1
|
|
maxLogEntries = 500
|
|
maxFailedLogLookupSources = 4096
|
|
maxConcurrentLogLookups = 32
|
|
maxHTTPHeaderBytes = 64 * 1024
|
|
maxPosterSize = 5 * 1024 * 1024 // 5MB
|
|
maxPosterStoreSize = int64(1 * 1024 * 1024 * 1024)
|
|
posterMaxAge = 3 * time.Hour
|
|
posterIDLength = 16
|
|
posterPerIPRateBurst = 3
|
|
posterPerIPRateSustained = 1
|
|
posterGlobalRateBurst = 8
|
|
posterGlobalRateSustained = 2
|
|
maxConcurrentPosterUploads = 4
|
|
posterUploadReadTimeout = 30 * time.Second
|
|
maxConnsPerIP = 5
|
|
maxGlobalConns = 100
|
|
maxRoomsPerIP = 3
|
|
maxRetainedRooms = 2000
|
|
connRateBurst = 5
|
|
connRateSustained = 1
|
|
reconnectTokenSize = 32
|
|
snapshotFormatVersion = 4
|
|
snapshotDebounce = 100 * time.Millisecond
|
|
snapshotFlushTimeout = 5 * time.Second
|
|
snapshotMaxFileSize = 4 * 1024 * 1024
|
|
)
|
|
|
|
var upgrader = websocket.Upgrader{
|
|
ReadBufferSize: 1024,
|
|
WriteBufferSize: 1024,
|
|
CheckOrigin: func(r *http.Request) bool { return true },
|
|
}
|
|
|
|
// --- Messages ---
|
|
|
|
type clientMsg struct {
|
|
Type string `json:"type"`
|
|
SessionID string `json:"sessionId,omitempty"`
|
|
PeerID string `json:"peerId,omitempty"`
|
|
ReconnectToken string `json:"reconnectToken,omitempty"`
|
|
ProtocolVersion int `json:"protocolVersion,omitempty"`
|
|
To string `json:"to,omitempty"`
|
|
Payload json.RawMessage `json:"payload,omitempty"`
|
|
}
|
|
|
|
type serverMsg struct {
|
|
Type string `json:"type"`
|
|
SessionID string `json:"sessionId,omitempty"`
|
|
PeerID string `json:"peerId,omitempty"`
|
|
HostPeerID string `json:"hostPeerId,omitempty"`
|
|
ReconnectToken string `json:"reconnectToken,omitempty"`
|
|
ProtocolVersion int `json:"protocolVersion,omitempty"`
|
|
From string `json:"from,omitempty"`
|
|
Peers []string `json:"peers,omitempty"`
|
|
Code string `json:"code,omitempty"`
|
|
Message string `json:"message,omitempty"`
|
|
Payload json.RawMessage `json:"payload,omitempty"`
|
|
}
|
|
|
|
// --- Client (serializes writes to a single goroutine) ---
|
|
|
|
type outboundFrame struct {
|
|
data []byte
|
|
written chan bool
|
|
}
|
|
|
|
type Client struct {
|
|
conn *websocket.Conn
|
|
send chan outboundFrame
|
|
done chan struct{}
|
|
closeOnce sync.Once
|
|
}
|
|
|
|
func newClient(conn *websocket.Conn) *Client {
|
|
c := &Client{conn: conn, send: make(chan outboundFrame, 64), done: make(chan struct{})}
|
|
go c.writePump()
|
|
return c
|
|
}
|
|
|
|
func (c *Client) writePump() {
|
|
ticker := time.NewTicker(pingInterval)
|
|
defer func() {
|
|
ticker.Stop()
|
|
c.close()
|
|
}()
|
|
for {
|
|
select {
|
|
case frame := <-c.send:
|
|
if err := c.conn.SetWriteDeadline(time.Now().Add(writeWait)); err != nil {
|
|
if frame.written != nil {
|
|
frame.written <- false
|
|
}
|
|
return
|
|
}
|
|
if err := c.conn.WriteMessage(websocket.TextMessage, frame.data); err != nil {
|
|
if frame.written != nil {
|
|
frame.written <- false
|
|
}
|
|
return
|
|
}
|
|
if frame.written != nil {
|
|
frame.written <- true
|
|
}
|
|
case <-c.done:
|
|
_ = c.conn.WriteMessage(websocket.CloseMessage, nil)
|
|
return
|
|
case <-ticker.C:
|
|
if err := c.conn.SetWriteDeadline(time.Now().Add(writeWait)); err != nil {
|
|
return
|
|
}
|
|
if err := c.conn.WriteMessage(websocket.PingMessage, nil); err != nil {
|
|
return
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
func (c *Client) enqueueFrame(frame outboundFrame) bool {
|
|
select {
|
|
case <-c.done:
|
|
return false
|
|
default:
|
|
}
|
|
|
|
select {
|
|
case <-c.done:
|
|
return false
|
|
case c.send <- frame:
|
|
return true
|
|
default:
|
|
c.close()
|
|
return false
|
|
}
|
|
}
|
|
|
|
func (c *Client) enqueue(data []byte) bool {
|
|
return c.enqueueFrame(outboundFrame{data: data})
|
|
}
|
|
|
|
func (c *Client) sendJSON(msg serverMsg) bool {
|
|
data, err := json.Marshal(msg)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
return c.enqueue(data)
|
|
}
|
|
|
|
func (c *Client) sendJSONAndWait(msg serverMsg) bool {
|
|
data, err := json.Marshal(msg)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
written := make(chan bool, 1)
|
|
if !c.enqueueFrame(outboundFrame{data: data, written: written}) {
|
|
return false
|
|
}
|
|
select {
|
|
case ok := <-written:
|
|
return ok
|
|
case <-time.After(writeWait):
|
|
return false
|
|
}
|
|
}
|
|
|
|
func (c *Client) close() {
|
|
c.closeOnce.Do(func() {
|
|
close(c.done)
|
|
_ = c.conn.Close()
|
|
})
|
|
}
|
|
|
|
// --- Room ---
|
|
|
|
type reconnectVerifier [sha256.Size]byte
|
|
|
|
type peerReservation struct {
|
|
verifier reconnectVerifier
|
|
absentSince time.Time
|
|
releasePending bool // runtime-only: pending releases are omitted from snapshots
|
|
releaseClient *Client // runtime-only: identifies the client that staged the release
|
|
}
|
|
|
|
type Room struct {
|
|
SessionID string
|
|
HostPeerID string
|
|
ProtocolVersion int
|
|
hostVerifier reconnectVerifier
|
|
peerReservations map[string]peerReservation
|
|
Peers map[string]*Client `json:"-"`
|
|
quotaOwnerKey string `json:"-"`
|
|
mu sync.RWMutex `json:"-"`
|
|
closing bool `json:"-"`
|
|
CreatedAt time.Time
|
|
LastActivityAt time.Time
|
|
}
|
|
|
|
// --- Snapshot types (on-disk JSON format) ---
|
|
|
|
// Nanosecond timestamps preserve exact absence state without the expansion of
|
|
// RFC3339 strings at the maximum admitted reservation count. Zero means the
|
|
// peer was connected when the snapshot was captured.
|
|
type peerReservationSnapshot struct {
|
|
Verifier string `json:"verifier"`
|
|
AbsentSinceUnixNano int64 `json:"absentSince,omitempty"`
|
|
}
|
|
|
|
type roomSnapshot struct {
|
|
SessionID string `json:"sessionId"`
|
|
HostPeerID string `json:"hostPeerId"`
|
|
ProtocolVersion int `json:"protocolVersion,omitempty"`
|
|
HostReconnectVerifier string `json:"hostReconnectVerifier"`
|
|
PeerReservations map[string]peerReservationSnapshot `json:"peerReservations,omitempty"`
|
|
PeerReconnectVerifiers map[string]string `json:"peerReconnectVerifiers,omitempty"` // v2/v3 decode only
|
|
CreatedAt time.Time `json:"createdAt"`
|
|
LastActivityAt time.Time `json:"lastActivityAt"`
|
|
}
|
|
|
|
type stateSnapshot struct {
|
|
Version int `json:"version"`
|
|
SavedAt time.Time `json:"savedAt"`
|
|
Rooms []roomSnapshot `json:"rooms"`
|
|
}
|
|
|
|
func mintReconnectToken() (string, reconnectVerifier, error) {
|
|
raw := make([]byte, reconnectTokenSize)
|
|
if _, err := rand.Read(raw); err != nil {
|
|
return "", reconnectVerifier{}, err
|
|
}
|
|
return base64.RawURLEncoding.EncodeToString(raw), sha256.Sum256(raw), nil
|
|
}
|
|
|
|
func reconnectVerifierFromToken(token string) (reconnectVerifier, bool) {
|
|
if len(token) != base64.RawURLEncoding.EncodedLen(reconnectTokenSize) {
|
|
return reconnectVerifier{}, false
|
|
}
|
|
raw, err := base64.RawURLEncoding.DecodeString(token)
|
|
if err != nil || len(raw) != reconnectTokenSize {
|
|
return reconnectVerifier{}, false
|
|
}
|
|
return sha256.Sum256(raw), true
|
|
}
|
|
|
|
func reconnectVerifierFromSnapshot(encoded string) (reconnectVerifier, bool) {
|
|
raw, err := base64.RawURLEncoding.DecodeString(encoded)
|
|
if err != nil || len(raw) != sha256.Size {
|
|
return reconnectVerifier{}, false
|
|
}
|
|
var verifier reconnectVerifier
|
|
copy(verifier[:], raw)
|
|
return verifier, true
|
|
}
|
|
|
|
func encodeReconnectVerifier(verifier reconnectVerifier) string {
|
|
return base64.RawURLEncoding.EncodeToString(verifier[:])
|
|
}
|
|
|
|
func reconnectVerifierMatches(expected, presented reconnectVerifier) bool {
|
|
return subtle.ConstantTimeCompare(expected[:], presented[:]) == 1
|
|
}
|
|
|
|
// pruneExpiredPeerReservationsLocked removes only expired, disconnected guest
|
|
// reservations. The caller must hold room.mu.
|
|
func pruneExpiredPeerReservationsLocked(room *Room, now time.Time) bool {
|
|
changed := false
|
|
for peerID, reservation := range room.peerReservations {
|
|
if reservation.releasePending ||
|
|
reservation.absentSince.IsZero() ||
|
|
now.Before(reservation.absentSince.Add(peerReservationGrace)) {
|
|
continue
|
|
}
|
|
if _, connected := room.Peers[peerID]; connected {
|
|
continue
|
|
}
|
|
delete(room.peerReservations, peerID)
|
|
changed = true
|
|
}
|
|
return changed
|
|
}
|
|
|
|
func (r *Room) peerIDs() []string {
|
|
ids := make([]string, 0, len(r.Peers))
|
|
for id := range r.Peers {
|
|
ids = append(ids, id)
|
|
}
|
|
return ids
|
|
}
|
|
|
|
func (r *Room) broadcastExcept(senderID string, msg serverMsg) {
|
|
data, err := json.Marshal(msg)
|
|
if err != nil {
|
|
return
|
|
}
|
|
// 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.Unlock()
|
|
|
|
for _, client := range targets {
|
|
client.enqueue(data)
|
|
}
|
|
}
|
|
|
|
func (r *Room) broadcastFrom(senderID string, sender *Client, msg serverMsg) bool {
|
|
data, err := json.Marshal(msg)
|
|
if err != nil {
|
|
return false
|
|
}
|
|
r.mu.Lock()
|
|
if r.Peers[senderID] != sender {
|
|
r.mu.Unlock()
|
|
return false
|
|
}
|
|
if r.closing {
|
|
r.mu.Unlock()
|
|
return true
|
|
}
|
|
targets := make([]*Client, 0, len(r.Peers)-1)
|
|
r.LastActivityAt = time.Now()
|
|
for id, client := range r.Peers {
|
|
if id != senderID {
|
|
targets = append(targets, client)
|
|
}
|
|
}
|
|
r.mu.Unlock()
|
|
|
|
for _, target := range targets {
|
|
target.enqueue(data)
|
|
}
|
|
return true
|
|
}
|
|
|
|
type directedSendResult uint8
|
|
|
|
const (
|
|
directedSenderUnavailable directedSendResult = iota
|
|
directedTargetMissing
|
|
directedTargetFound
|
|
directedSendSuppressed
|
|
)
|
|
|
|
func (r *Room) sendFrom(senderID string, sender *Client, targetID string, msg serverMsg) directedSendResult {
|
|
data, err := json.Marshal(msg)
|
|
if err != nil {
|
|
return directedSenderUnavailable
|
|
}
|
|
r.mu.Lock()
|
|
if r.Peers[senderID] != sender {
|
|
r.mu.Unlock()
|
|
return directedSenderUnavailable
|
|
}
|
|
if r.closing {
|
|
r.mu.Unlock()
|
|
return directedSendSuppressed
|
|
}
|
|
target, ok := r.Peers[targetID]
|
|
if ok {
|
|
r.LastActivityAt = time.Now()
|
|
}
|
|
r.mu.Unlock()
|
|
if !ok {
|
|
return directedTargetMissing
|
|
}
|
|
target.enqueue(data)
|
|
return directedTargetFound
|
|
}
|
|
|
|
// --- Log store ---
|
|
|
|
const logFileExt = ".log"
|
|
|
|
var errLogStoreFull = errors.New("log store full")
|
|
|
|
// logStore keeps diagnostic uploads capped by artifact count; a full store
|
|
// rejects new uploads rather than evicting logs someone may still be reading.
|
|
type logStore struct {
|
|
artifactStore
|
|
rateLimit map[string]time.Time // IP -> last upload time
|
|
failedLookupRate map[string]*rateLimiter
|
|
}
|
|
|
|
func newLogStore(dir string) *logStore {
|
|
return newLogStoreWithRemover(dir, os.Remove)
|
|
}
|
|
|
|
func newLogStoreWithRemover(dir string, removeFile func(string) error) *logStore {
|
|
if err := os.MkdirAll(dir, 0755); err != nil {
|
|
log.Fatalf("failed to create log dir %s: %v", dir, err)
|
|
}
|
|
ls := &logStore{
|
|
artifactStore: artifactStore{
|
|
entries: make(map[string]artifactEntry),
|
|
pendingRemovals: make(map[string]pendingRemoval),
|
|
dir: dir,
|
|
name: "logs",
|
|
maxAge: logMaxAge,
|
|
removeFile: removeFile,
|
|
generateID: generateLogID,
|
|
idFromFilename: logIDFromFilename,
|
|
acceptLoaded: func(_ string, size int64) (string, bool) {
|
|
return "", size > 0 && size <= maxLogSize
|
|
},
|
|
limit: maxLogEntries,
|
|
cost: func(int64) int64 { return 1 },
|
|
pendingCost: func(pendingRemoval) int64 { return 1 },
|
|
errFull: errLogStoreFull,
|
|
},
|
|
rateLimit: make(map[string]time.Time),
|
|
failedLookupRate: make(map[string]*rateLimiter),
|
|
}
|
|
ls.startupErr = ls.loadExisting(time.Now())
|
|
return ls
|
|
}
|
|
|
|
func (ls *logStore) filePath(id string) string {
|
|
return ls.artifactStore.filePath(id + logFileExt)
|
|
}
|
|
|
|
func generateLogID() string {
|
|
return generateID(logIDLength)
|
|
}
|
|
|
|
func logIDFromFilename(filename string) (string, bool) {
|
|
if filepath.Ext(filename) != logFileExt {
|
|
return "", false
|
|
}
|
|
id := strings.TrimSuffix(filename, logFileExt)
|
|
return id, validID(id, logIDLength)
|
|
}
|
|
|
|
func (ls *logStore) store(data []byte, now time.Time) (string, artifactEntry, error) {
|
|
if len(data) == 0 {
|
|
return "", artifactEntry{}, errors.New("empty log")
|
|
}
|
|
if len(data) > maxLogSize {
|
|
return "", artifactEntry{}, errors.New("log too large")
|
|
}
|
|
return ls.put(data, logFileExt, "", now)
|
|
}
|
|
|
|
func (ls *logStore) lookup(id string, now time.Time) (artifactEntry, bool, error) {
|
|
if !validID(id, logIDLength) {
|
|
return artifactEntry{}, false, nil
|
|
}
|
|
return ls.lookupEntry(id, now, nil)
|
|
}
|
|
|
|
func (ls *logStore) allowFailedLookup(source string, now time.Time) bool {
|
|
ls.mu.Lock()
|
|
defer ls.mu.Unlock()
|
|
limiter := ls.failedLookupRate[source]
|
|
if limiter == nil {
|
|
cleanupRateLimiters(ls.failedLookupRate, now, nil)
|
|
if len(ls.failedLookupRate) >= maxFailedLogLookupSources {
|
|
return false
|
|
}
|
|
limiter = newRateLimiterAt(logLookupRateBurst, logLookupRateSustained, now)
|
|
ls.failedLookupRate[source] = limiter
|
|
}
|
|
return limiter.allowAt(now)
|
|
}
|
|
|
|
func (ls *logStore) cleanup(now time.Time) error {
|
|
ls.mu.Lock()
|
|
defer ls.mu.Unlock()
|
|
removalErr := ls.cleanupLocked(now)
|
|
cleanupRateWindows(ls.rateLimit, now, logRateInterval)
|
|
cleanupRateLimiters(ls.failedLookupRate, now, nil)
|
|
return removalErr
|
|
}
|
|
|
|
// --- Poster store ---
|
|
|
|
var errPosterStoreFull = errors.New("poster store full")
|
|
|
|
// posterStore caps shared posters by accounted bytes and evicts the oldest to
|
|
// admit a new upload.
|
|
type posterStore struct {
|
|
artifactStore
|
|
}
|
|
|
|
func newPosterStore(dir string, maxBytes int64, maxAge time.Duration) *posterStore {
|
|
return newPosterStoreWithRemover(dir, maxBytes, maxAge, os.Remove)
|
|
}
|
|
|
|
func newPosterStoreWithRemover(
|
|
dir string,
|
|
maxBytes int64,
|
|
maxAge time.Duration,
|
|
removeFile func(string) error,
|
|
) *posterStore {
|
|
if err := os.MkdirAll(dir, 0755); err != nil {
|
|
log.Fatalf("failed to create poster dir %s: %v", dir, err)
|
|
}
|
|
ps := &posterStore{artifactStore{
|
|
entries: make(map[string]artifactEntry),
|
|
pendingRemovals: make(map[string]pendingRemoval),
|
|
dir: dir,
|
|
name: "posters",
|
|
maxAge: maxAge,
|
|
removeFile: removeFile,
|
|
generateID: generatePosterID,
|
|
idFromFilename: posterIDFromFilename,
|
|
acceptLoaded: func(filename string, _ int64) (string, bool) {
|
|
return posterContentTypeForExt(filepath.Ext(filename))
|
|
},
|
|
limit: maxBytes,
|
|
cost: func(size int64) int64 { return size },
|
|
pendingCost: func(pending pendingRemoval) int64 {
|
|
// Unknown debt cannot be sized safely, so it is kept out of the
|
|
// quota: a permanent directory or stat failure must not deny
|
|
// otherwise capacity-safe uploads.
|
|
if !pending.sizeKnown {
|
|
return 0
|
|
}
|
|
return pending.size
|
|
},
|
|
evictToFit: true,
|
|
retryKnownDebtOnPut: true,
|
|
errFull: errPosterStoreFull,
|
|
}}
|
|
ps.startupErr = ps.loadExisting(time.Now())
|
|
return ps
|
|
}
|
|
|
|
func generatePosterID() string {
|
|
return generateID(posterIDLength)
|
|
}
|
|
|
|
func posterExtForContentType(contentType string) (string, bool) {
|
|
switch strings.ToLower(strings.SplitN(contentType, ";", 2)[0]) {
|
|
case "image/jpeg":
|
|
return ".jpg", true
|
|
case "image/png":
|
|
return ".png", true
|
|
case "image/gif":
|
|
return ".gif", true
|
|
case "image/webp":
|
|
return ".webp", true
|
|
default:
|
|
return "", false
|
|
}
|
|
}
|
|
|
|
func posterContentTypeForExt(ext string) (string, bool) {
|
|
switch strings.ToLower(ext) {
|
|
case ".jpg", ".jpeg":
|
|
return "image/jpeg", true
|
|
case ".png":
|
|
return "image/png", true
|
|
case ".gif":
|
|
return "image/gif", true
|
|
case ".webp":
|
|
return "image/webp", true
|
|
default:
|
|
return "", false
|
|
}
|
|
}
|
|
|
|
func posterIDFromFilename(filename string) (string, bool) {
|
|
if filename == "" || strings.ContainsAny(filename, `/\\`) {
|
|
return "", false
|
|
}
|
|
ext := filepath.Ext(filename)
|
|
if _, ok := posterContentTypeForExt(ext); !ok {
|
|
return "", false
|
|
}
|
|
id := strings.TrimSuffix(filename, ext)
|
|
if !validID(id, posterIDLength) {
|
|
return "", false
|
|
}
|
|
return id, true
|
|
}
|
|
|
|
func (ps *posterStore) store(data []byte, contentType string, now time.Time) (string, artifactEntry, error) {
|
|
entrySize := int64(len(data))
|
|
if entrySize <= 0 {
|
|
return "", artifactEntry{}, errors.New("empty poster")
|
|
}
|
|
if entrySize > ps.limit {
|
|
return "", artifactEntry{}, errors.New("poster exceeds store size")
|
|
}
|
|
ext, ok := posterExtForContentType(contentType)
|
|
if !ok {
|
|
return "", artifactEntry{}, errors.New("unsupported poster type")
|
|
}
|
|
return ps.put(data, ext, strings.ToLower(strings.SplitN(contentType, ";", 2)[0]), now)
|
|
}
|
|
|
|
func (ps *posterStore) lookup(filename string, now time.Time) (artifactEntry, bool, error) {
|
|
id, ok := posterIDFromFilename(filename)
|
|
if !ok {
|
|
return artifactEntry{}, false, nil
|
|
}
|
|
return ps.lookupEntry(id, now, func(entry artifactEntry) bool {
|
|
return entry.Filename == filename
|
|
})
|
|
}
|
|
|
|
// --- Snapshotter (single-writer, debounced, atomic disk persistence) ---
|
|
|
|
var errSnapshotterStopped = errors.New("snapshot writer is stopped")
|
|
|
|
type terminalMutationOutcome struct {
|
|
err error
|
|
deliver bool
|
|
}
|
|
|
|
type terminalMutationTicket struct {
|
|
seq uint64
|
|
result <-chan terminalMutationOutcome
|
|
}
|
|
|
|
type registeredTerminalMutation struct {
|
|
seq uint64
|
|
complete func(error) terminalMutationOutcome
|
|
result chan terminalMutationOutcome
|
|
}
|
|
|
|
type snapshotter struct {
|
|
path string
|
|
dir string
|
|
trigger chan struct{}
|
|
urgent chan struct{}
|
|
flush chan chan error
|
|
exited chan struct{}
|
|
build func() stateSnapshot
|
|
capture func(func() uint64) (stateSnapshot, uint64)
|
|
persist func([]byte) error
|
|
syncDir func(string) error
|
|
debounce time.Duration
|
|
writeMu sync.Mutex
|
|
beforeDebounceWait func() // test-only signal after a trigger enters its debounce window
|
|
beforeCapture func() // test-only barrier immediately before generation capture
|
|
afterSequenceCapture func() // test-only barrier inside the protected capture boundary
|
|
|
|
stateMu sync.Mutex
|
|
dirtySeq uint64
|
|
durableSeq uint64
|
|
terminals []*registeredTerminalMutation
|
|
stopped bool
|
|
|
|
stopOnce sync.Once
|
|
stopErr error
|
|
|
|
errMu sync.Mutex
|
|
lastErrLog time.Time
|
|
|
|
dirErrMu sync.Mutex
|
|
lastDirErrLog time.Time
|
|
}
|
|
|
|
func newSnapshotter(path string, build func() stateSnapshot) *snapshotter {
|
|
dir := filepath.Dir(path)
|
|
if err := os.MkdirAll(dir, 0755); err != nil {
|
|
log.Printf("snapshot: mkdir %s: %v", dir, err)
|
|
}
|
|
sn := &snapshotter{
|
|
path: path,
|
|
dir: dir,
|
|
trigger: make(chan struct{}, 1),
|
|
urgent: make(chan struct{}, 1),
|
|
flush: make(chan chan error),
|
|
exited: make(chan struct{}),
|
|
build: build,
|
|
syncDir: syncSnapshotDirectory,
|
|
debounce: snapshotDebounce,
|
|
}
|
|
sn.capture = func(captureSequence func() uint64) (stateSnapshot, uint64) {
|
|
targetSeq := captureSequence()
|
|
return sn.build(), targetSeq
|
|
}
|
|
sn.persist = sn.persistAtomic
|
|
return sn
|
|
}
|
|
|
|
// recordMutation publishes a protected identity, membership, or reservation
|
|
// mutation to the single writer. Callers record after changing state and before
|
|
// releasing the lock that made the mutation visible.
|
|
func (sn *snapshotter) recordMutation() uint64 {
|
|
sn.stateMu.Lock()
|
|
if sn.stopped {
|
|
sn.stateMu.Unlock()
|
|
return 0
|
|
}
|
|
sn.dirtySeq++
|
|
seq := sn.dirtySeq
|
|
sn.signalLocked()
|
|
sn.stateMu.Unlock()
|
|
return seq
|
|
}
|
|
|
|
// recordTerminalMutation atomically publishes a protected mutation together
|
|
// with its outcome channel. A buffered result retains even an immediate write
|
|
// failure until the handler begins waiting.
|
|
func (sn *snapshotter) recordTerminalMutation(
|
|
complete func(error) terminalMutationOutcome,
|
|
) *terminalMutationTicket {
|
|
result := make(chan terminalMutationOutcome, 1)
|
|
sn.stateMu.Lock()
|
|
if sn.stopped {
|
|
sn.stateMu.Unlock()
|
|
if complete == nil {
|
|
result <- terminalMutationOutcome{err: errSnapshotterStopped, deliver: true}
|
|
} else {
|
|
// The caller still holds the protected state lock. Run the
|
|
// rollback continuation asynchronously so it can acquire the
|
|
// normal s.mu -> room.mu order after the caller unlocks.
|
|
go func() {
|
|
result <- complete(errSnapshotterStopped)
|
|
}()
|
|
}
|
|
return &terminalMutationTicket{result: result}
|
|
}
|
|
sn.dirtySeq++
|
|
seq := sn.dirtySeq
|
|
sn.terminals = append(sn.terminals, ®isteredTerminalMutation{
|
|
seq: seq,
|
|
complete: complete,
|
|
result: result,
|
|
})
|
|
sn.signalLocked()
|
|
select {
|
|
case sn.urgent <- struct{}{}:
|
|
default:
|
|
}
|
|
sn.stateMu.Unlock()
|
|
return &terminalMutationTicket{seq: seq, result: result}
|
|
}
|
|
|
|
func (sn *snapshotter) signalLocked() {
|
|
select {
|
|
case sn.trigger <- struct{}{}:
|
|
default:
|
|
}
|
|
}
|
|
|
|
func (sn *snapshotter) captureSequence() uint64 {
|
|
sn.stateMu.Lock()
|
|
targetSeq := sn.dirtySeq
|
|
afterCapture := sn.afterSequenceCapture
|
|
sn.stateMu.Unlock()
|
|
if afterCapture != nil {
|
|
afterCapture()
|
|
}
|
|
return targetSeq
|
|
}
|
|
|
|
func (sn *snapshotter) waitForDurable(ticket *terminalMutationTicket) terminalMutationOutcome {
|
|
return <-ticket.result
|
|
}
|
|
|
|
func (sn *snapshotter) waitForDurableWithin(
|
|
ticket *terminalMutationTicket,
|
|
timeout time.Duration,
|
|
) terminalMutationOutcome {
|
|
timer := time.NewTimer(timeout)
|
|
defer timer.Stop()
|
|
select {
|
|
case outcome := <-ticket.result:
|
|
return outcome
|
|
case <-timer.C:
|
|
return terminalMutationOutcome{
|
|
err: errors.New("snapshot rewrite: timed out waiting for persistence"),
|
|
deliver: true,
|
|
}
|
|
}
|
|
}
|
|
|
|
func (sn *snapshotter) run() {
|
|
defer close(sn.exited)
|
|
for {
|
|
select {
|
|
case <-sn.trigger:
|
|
if sn.beforeDebounceWait != nil {
|
|
sn.beforeDebounceWait()
|
|
}
|
|
timer := time.NewTimer(sn.debounce)
|
|
select {
|
|
case <-timer.C:
|
|
case <-sn.urgent:
|
|
case reply := <-sn.flush:
|
|
if !timer.Stop() {
|
|
select {
|
|
case <-timer.C:
|
|
default:
|
|
}
|
|
}
|
|
err := sn.flushLatestAndStop()
|
|
reply <- err
|
|
return
|
|
}
|
|
if !timer.Stop() {
|
|
select {
|
|
case <-timer.C:
|
|
default:
|
|
}
|
|
}
|
|
// Drain tokens queued before capture. A mutation recorded after
|
|
// capture re-arms the channels and therefore requires a later write.
|
|
select {
|
|
case <-sn.trigger:
|
|
default:
|
|
}
|
|
select {
|
|
case <-sn.urgent:
|
|
default:
|
|
}
|
|
_, err := sn.writeNextGeneration()
|
|
if err != nil {
|
|
sn.logWriteErr(err)
|
|
}
|
|
case reply := <-sn.flush:
|
|
err := sn.flushLatestAndStop()
|
|
reply <- err
|
|
return
|
|
}
|
|
}
|
|
}
|
|
|
|
// write is the narrowly serialized storage entry retained for atomic-storage
|
|
// tests. Production mutations use writeNextGeneration so generation outcomes
|
|
// cannot bypass the single writer.
|
|
func (sn *snapshotter) write() error {
|
|
sn.writeMu.Lock()
|
|
defer sn.writeMu.Unlock()
|
|
_, err := sn.captureAndPersist()
|
|
return err
|
|
}
|
|
|
|
func (sn *snapshotter) captureAndPersist() (uint64, error) {
|
|
snapshot, targetSeq := sn.capture(sn.captureSequence)
|
|
data, err := json.Marshal(snapshot)
|
|
if err != nil {
|
|
return targetSeq, err
|
|
}
|
|
if len(data) > snapshotMaxFileSize {
|
|
return targetSeq, fmt.Errorf("snapshot exceeds maximum size: %d > %d bytes", len(data), snapshotMaxFileSize)
|
|
}
|
|
return targetSeq, sn.persist(data)
|
|
}
|
|
|
|
func (sn *snapshotter) writeNextGeneration() (bool, error) {
|
|
sn.stateMu.Lock()
|
|
if sn.dirtySeq <= sn.durableSeq {
|
|
sn.stateMu.Unlock()
|
|
return false, nil
|
|
}
|
|
beforeCapture := sn.beforeCapture
|
|
sn.stateMu.Unlock()
|
|
if beforeCapture != nil {
|
|
beforeCapture()
|
|
}
|
|
sn.writeMu.Lock()
|
|
targetSeq, err := sn.captureAndPersist()
|
|
sn.writeMu.Unlock()
|
|
|
|
sn.stateMu.Lock()
|
|
if err == nil && targetSeq > sn.durableSeq {
|
|
sn.durableSeq = targetSeq
|
|
}
|
|
coveredCount := 0
|
|
for coveredCount < len(sn.terminals) && sn.terminals[coveredCount].seq <= targetSeq {
|
|
coveredCount++
|
|
}
|
|
covered := sn.terminals[:coveredCount:coveredCount]
|
|
sn.terminals = sn.terminals[coveredCount:]
|
|
if len(sn.terminals) == 0 {
|
|
sn.terminals = nil
|
|
}
|
|
sn.stateMu.Unlock()
|
|
|
|
// Continuations are part of the writer barrier. In particular, a failed
|
|
// staged release rolls back and records its corrective generation before
|
|
// this writer can capture any queued later mutation.
|
|
for _, terminal := range covered {
|
|
outcome := terminalMutationOutcome{err: err, deliver: true}
|
|
if terminal.complete != nil {
|
|
outcome = terminal.complete(err)
|
|
}
|
|
terminal.result <- outcome
|
|
}
|
|
return true, err
|
|
}
|
|
|
|
func (sn *snapshotter) flushLatestAndStop() error {
|
|
for {
|
|
sn.stateMu.Lock()
|
|
if sn.dirtySeq <= sn.durableSeq {
|
|
sn.stopped = true
|
|
sn.stateMu.Unlock()
|
|
return nil
|
|
}
|
|
sn.stateMu.Unlock()
|
|
|
|
_, err := sn.writeNextGeneration()
|
|
if err == nil {
|
|
continue
|
|
}
|
|
sn.logWriteErr(err)
|
|
sn.failPendingAndStop(err)
|
|
return err
|
|
}
|
|
}
|
|
|
|
func (sn *snapshotter) failPendingAndStop(err error) {
|
|
sn.stateMu.Lock()
|
|
sn.stopped = true
|
|
pending := sn.terminals
|
|
sn.terminals = nil
|
|
sn.stateMu.Unlock()
|
|
for _, terminal := range pending {
|
|
outcome := terminalMutationOutcome{err: err, deliver: true}
|
|
if terminal.complete != nil {
|
|
outcome = terminal.complete(err)
|
|
}
|
|
terminal.result <- outcome
|
|
}
|
|
}
|
|
|
|
func syncSnapshotDirectory(dir string) error {
|
|
d, err := os.Open(dir)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
syncErr := d.Sync()
|
|
closeErr := d.Close()
|
|
return errors.Join(syncErr, closeErr)
|
|
}
|
|
|
|
func (sn *snapshotter) persistAtomic(data []byte) error {
|
|
tmpPath := sn.path + ".tmp"
|
|
f, err := os.OpenFile(tmpPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0644)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if _, err := f.Write(data); err != nil {
|
|
f.Close()
|
|
os.Remove(tmpPath)
|
|
return err
|
|
}
|
|
if err := f.Sync(); err != nil {
|
|
f.Close()
|
|
os.Remove(tmpPath)
|
|
return err
|
|
}
|
|
if err := f.Close(); err != nil {
|
|
os.Remove(tmpPath)
|
|
return err
|
|
}
|
|
if err := os.Rename(tmpPath, sn.path); err != nil {
|
|
os.Remove(tmpPath)
|
|
return err
|
|
}
|
|
// Rename is the commit boundary: the replacement is file-synced and
|
|
// non-torn. Parent-directory sync adds crash durability where supported,
|
|
// but its post-commit failure must not report the mutation as uncommitted.
|
|
if err := sn.syncDir(sn.dir); err != nil {
|
|
sn.logDirSyncErr(err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func (sn *snapshotter) flushAndStop(timeout time.Duration) error {
|
|
sn.stopOnce.Do(func() {
|
|
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
|
defer cancel()
|
|
reply := make(chan error, 1)
|
|
select {
|
|
case sn.flush <- reply:
|
|
case <-ctx.Done():
|
|
sn.stopErr = errors.New("snapshot flush: timed out sending flush signal")
|
|
return
|
|
}
|
|
select {
|
|
case sn.stopErr = <-reply:
|
|
case <-ctx.Done():
|
|
sn.stopErr = errors.New("snapshot flush: timed out waiting for write")
|
|
return
|
|
}
|
|
select {
|
|
case <-sn.exited:
|
|
case <-ctx.Done():
|
|
}
|
|
})
|
|
return sn.stopErr
|
|
}
|
|
|
|
// logWriteErr throttles pre-commit snapshot-write error spam to once per hour.
|
|
func (sn *snapshotter) logWriteErr(err error) {
|
|
sn.errMu.Lock()
|
|
defer sn.errMu.Unlock()
|
|
if time.Since(sn.lastErrLog) < time.Hour {
|
|
return
|
|
}
|
|
sn.lastErrLog = time.Now()
|
|
log.Printf("snapshot: write failed before rename commit: %v", err)
|
|
}
|
|
|
|
// logDirSyncErr has an independent throttle so a degraded post-rename warning
|
|
// cannot suppress a later pre-commit persistence error.
|
|
func (sn *snapshotter) logDirSyncErr(err error) {
|
|
sn.dirErrMu.Lock()
|
|
defer sn.dirErrMu.Unlock()
|
|
if time.Since(sn.lastDirErrLog) < time.Hour {
|
|
return
|
|
}
|
|
sn.lastDirErrLog = time.Now()
|
|
log.Printf("snapshot: parent directory sync failed after rename commit: %v", err)
|
|
}
|
|
|
|
// --- Server ---
|
|
type removalErrorThrottle struct {
|
|
mu sync.Mutex
|
|
lastLog map[string]time.Time
|
|
}
|
|
|
|
func (s *Server) logRemovalError(store, operation string, err error) {
|
|
key := store + ":" + operation
|
|
s.removalErrors.mu.Lock()
|
|
defer s.removalErrors.mu.Unlock()
|
|
if last := s.removalErrors.lastLog[key]; !last.IsZero() && time.Since(last) < time.Hour {
|
|
return
|
|
}
|
|
if s.removalErrors.lastLog == nil {
|
|
s.removalErrors.lastLog = make(map[string]time.Time)
|
|
}
|
|
s.removalErrors.lastLog[key] = time.Now()
|
|
category := "other"
|
|
switch {
|
|
case errors.Is(err, errArtifactOutsideStore):
|
|
category = "confinement"
|
|
case errors.Is(err, fs.ErrPermission):
|
|
category = "permission"
|
|
case errors.Is(err, syscall.ENOSPC):
|
|
category = "capacity"
|
|
case errors.Is(err, syscall.EROFS):
|
|
category = "read_only"
|
|
case errors.Is(err, syscall.EBUSY):
|
|
category = "busy"
|
|
case errors.Is(err, syscall.ENOTEMPTY), errors.Is(err, syscall.EEXIST):
|
|
category = "not_empty"
|
|
}
|
|
var errno syscall.Errno
|
|
if errors.As(err, &errno) {
|
|
log.Printf("%s: %s removal failed: category=%s errno=%d", store, operation, category, errno)
|
|
return
|
|
}
|
|
log.Printf("%s: %s removal failed: category=%s errno=unknown", store, operation, category)
|
|
}
|
|
|
|
type Server struct {
|
|
rooms map[string]*Room
|
|
logs *logStore
|
|
posters *posterStore
|
|
posterUploads *posterUploadLimiter
|
|
logLookups chan struct{}
|
|
posterBodyReadTimeout time.Duration
|
|
conns *connTracker
|
|
clientIPs clientIPResolver
|
|
snap *snapshotter
|
|
oauth *oauthProxy // nil when OAUTH_BASE_URL is unset
|
|
removalErrors removalErrorThrottle
|
|
beforeJoinRoomLock func() // test-only deterministic admission barrier
|
|
beforeLeaveRoomLock func() // test-only mutation/capture ordering barrier
|
|
beforeTerminalDelivery func() // test-only post-persistence, pre-delivery barrier
|
|
mu sync.RWMutex
|
|
}
|
|
|
|
func newServer(logDir, stateFile, posterDir string, clientIPs clientIPResolver) *Server {
|
|
s := &Server{
|
|
rooms: make(map[string]*Room),
|
|
logs: newLogStore(logDir),
|
|
posters: newPosterStore(posterDir, maxPosterStoreSize, posterMaxAge),
|
|
posterUploads: newPosterUploadLimiter(posterPerIPRateBurst, posterPerIPRateSustained, posterGlobalRateBurst, posterGlobalRateSustained, maxConcurrentPosterUploads, time.Now()),
|
|
logLookups: make(chan struct{}, maxConcurrentLogLookups),
|
|
posterBodyReadTimeout: posterUploadReadTimeout,
|
|
conns: newConnTracker(),
|
|
clientIPs: clientIPs,
|
|
}
|
|
if s.logs.startupErr != nil {
|
|
s.logRemovalError("logs", "startup", s.logs.startupErr)
|
|
}
|
|
if s.posters.startupErr != nil {
|
|
s.logRemovalError("posters", "startup", s.posters.startupErr)
|
|
}
|
|
if p, ok := oauthConfigFromEnv(clientIPs); ok {
|
|
s.oauth = p
|
|
log.Printf("oauth: proxy enabled (base=%s, services=%d)", p.baseURL, len(p.services))
|
|
}
|
|
s.snap = newSnapshotter(stateFile, s.buildSnapshot)
|
|
s.snap.capture = s.captureSnapshot
|
|
rewriteReservations, err := s.loadSnapshot(stateFile)
|
|
if err != nil {
|
|
log.Printf("snapshot: load error: %v", err)
|
|
}
|
|
go s.snap.run()
|
|
if rewriteReservations {
|
|
ticket := s.snap.recordTerminalMutation(nil)
|
|
if outcome := s.snap.waitForDurableWithin(ticket, snapshotFlushTimeout); outcome.err != nil {
|
|
log.Printf("snapshot: v4 reservation rewrite failed: %v", outcome.err)
|
|
}
|
|
}
|
|
go s.cleanupLoop()
|
|
return s
|
|
}
|
|
|
|
// removeRoomLocked removes room only while it is still the authoritative map
|
|
// entry. The caller must hold s.mu. A current-process quota reservation follows
|
|
// the retained room and is returned exactly once by successful removal.
|
|
func (s *Server) removeRoomLocked(sessionID string, room *Room) bool {
|
|
if s.rooms[sessionID] != room {
|
|
return false
|
|
}
|
|
delete(s.rooms, sessionID)
|
|
if room.quotaOwnerKey != "" {
|
|
s.conns.releaseRoom(room.quotaOwnerKey)
|
|
}
|
|
return true
|
|
}
|
|
|
|
// buildSnapshot is the synchronous storage-test entry. Production capture uses
|
|
// captureSnapshot so the copied state and its covered generation share one
|
|
// ordering boundary.
|
|
func (s *Server) buildSnapshot() stateSnapshot {
|
|
snapshot, _ := s.captureSnapshot(func() uint64 { return 0 })
|
|
return snapshot
|
|
}
|
|
|
|
// captureSnapshot freezes every durable room mutation under the established
|
|
// s.mu -> room.mu order, then captures the covered sequence while those locks
|
|
// remain held. A mutation is therefore either both present and covered, or
|
|
// neither present nor covered. Locks are released before marshal or disk I/O.
|
|
func (s *Server) captureSnapshot(captureSequence func() uint64) (stateSnapshot, uint64) {
|
|
s.mu.RLock()
|
|
rooms := make([]*Room, 0, len(s.rooms))
|
|
for _, room := range s.rooms {
|
|
room.mu.RLock()
|
|
rooms = append(rooms, room)
|
|
}
|
|
|
|
targetSeq := captureSequence()
|
|
snapshot := stateSnapshot{
|
|
Version: snapshotFormatVersion,
|
|
SavedAt: time.Now(),
|
|
Rooms: make([]roomSnapshot, 0, len(rooms)),
|
|
}
|
|
for _, room := range rooms {
|
|
var reservations map[string]peerReservationSnapshot
|
|
if len(room.peerReservations) != 0 {
|
|
for peerID, reservation := range room.peerReservations {
|
|
if reservation.releasePending {
|
|
continue
|
|
}
|
|
if reservations == nil {
|
|
reservations = make(map[string]peerReservationSnapshot, len(room.peerReservations))
|
|
}
|
|
var absentSinceUnixNano int64
|
|
if !reservation.absentSince.IsZero() {
|
|
absentSinceUnixNano = reservation.absentSince.UnixNano()
|
|
}
|
|
reservations[peerID] = peerReservationSnapshot{
|
|
Verifier: encodeReconnectVerifier(reservation.verifier),
|
|
AbsentSinceUnixNano: absentSinceUnixNano,
|
|
}
|
|
}
|
|
}
|
|
snapshot.Rooms = append(snapshot.Rooms, roomSnapshot{
|
|
SessionID: room.SessionID,
|
|
HostPeerID: room.HostPeerID,
|
|
ProtocolVersion: room.ProtocolVersion,
|
|
HostReconnectVerifier: encodeReconnectVerifier(room.hostVerifier),
|
|
PeerReservations: reservations,
|
|
CreatedAt: room.CreatedAt,
|
|
LastActivityAt: room.LastActivityAt,
|
|
})
|
|
}
|
|
|
|
for index := len(rooms) - 1; index >= 0; index-- {
|
|
rooms[index].mu.RUnlock()
|
|
}
|
|
s.mu.RUnlock()
|
|
return snapshot, targetSeq
|
|
}
|
|
|
|
// loadSnapshot restores rooms from disk on startup. The returned rewrite flag
|
|
// reports reservation migration, initialization, or pruning that must be
|
|
// persisted before serving. Missing/corrupt files still allow startup.
|
|
func (s *Server) loadSnapshot(path string) (bool, error) {
|
|
data, err := os.ReadFile(path)
|
|
if err != nil {
|
|
if errors.Is(err, fs.ErrNotExist) {
|
|
log.Printf("snapshot: no file at %s, starting fresh", path)
|
|
return false, nil
|
|
}
|
|
log.Printf("snapshot: read error, starting fresh: %v", err)
|
|
return false, nil
|
|
}
|
|
if len(data) > snapshotMaxFileSize {
|
|
log.Printf("snapshot: file too large (%d bytes), starting fresh", len(data))
|
|
return false, nil
|
|
}
|
|
var snap stateSnapshot
|
|
if err := json.Unmarshal(data, &snap); err != nil {
|
|
log.Printf("snapshot: corrupt file at %s, starting fresh: %v", path, err)
|
|
return false, nil
|
|
}
|
|
if snap.Version != 2 && snap.Version != 3 && snap.Version != snapshotFormatVersion {
|
|
log.Printf("snapshot: unknown version %d, starting fresh", snap.Version)
|
|
return false, nil
|
|
}
|
|
now := time.Now()
|
|
rewriteReservations := snap.Version != snapshotFormatVersion
|
|
restored := make(map[string]*Room, min(len(snap.Rooms), maxRetainedRooms))
|
|
skipped := 0
|
|
for _, r := range snap.Rooms {
|
|
if !validRelayID(r.SessionID, maxSessionIDLength) || !validRelayID(r.HostPeerID, maxPeerIDLength) {
|
|
skipped++
|
|
continue
|
|
}
|
|
hostVerifier, ok := reconnectVerifierFromSnapshot(r.HostReconnectVerifier)
|
|
if !ok {
|
|
skipped++
|
|
continue
|
|
}
|
|
if r.ProtocolVersion != legacyRelayProtocolVersion && r.ProtocolVersion != relayProtocolVersion {
|
|
skipped++
|
|
continue
|
|
}
|
|
if now.Sub(r.CreatedAt) > roomMaxAge || now.Sub(r.LastActivityAt) > emptyRoomMaxAge {
|
|
skipped++
|
|
continue
|
|
}
|
|
if _, duplicate := restored[r.SessionID]; duplicate {
|
|
skipped++
|
|
continue
|
|
}
|
|
if len(restored) >= maxRetainedRooms {
|
|
log.Printf("snapshot: too many retained rooms, starting fresh")
|
|
return false, nil
|
|
}
|
|
|
|
var reservations map[string]peerReservation
|
|
validReservations := true
|
|
switch {
|
|
case r.ProtocolVersion == legacyRelayProtocolVersion &&
|
|
(len(r.PeerReservations) != 0 || len(r.PeerReconnectVerifiers) != 0):
|
|
validReservations = false
|
|
case snap.Version == snapshotFormatVersion && len(r.PeerReconnectVerifiers) != 0:
|
|
validReservations = false
|
|
case snap.Version == snapshotFormatVersion:
|
|
if len(r.PeerReservations) > maxRoomSize-1 {
|
|
validReservations = false
|
|
break
|
|
}
|
|
if len(r.PeerReservations) != 0 {
|
|
reservations = make(map[string]peerReservation, len(r.PeerReservations))
|
|
}
|
|
for peerID, encoded := range r.PeerReservations {
|
|
verifier, verifierOK := reconnectVerifierFromSnapshot(encoded.Verifier)
|
|
if !validRelayID(peerID, maxPeerIDLength) ||
|
|
peerID == r.HostPeerID ||
|
|
!verifierOK {
|
|
validReservations = false
|
|
break
|
|
}
|
|
absentSince := now
|
|
if encoded.AbsentSinceUnixNano == 0 {
|
|
rewriteReservations = true
|
|
} else {
|
|
absentSince = time.Unix(0, encoded.AbsentSinceUnixNano)
|
|
if absentSince.After(now) {
|
|
absentSince = now
|
|
rewriteReservations = true
|
|
} else if !now.Before(absentSince.Add(peerReservationGrace)) {
|
|
rewriteReservations = true
|
|
continue
|
|
}
|
|
}
|
|
reservations[peerID] = peerReservation{
|
|
verifier: verifier,
|
|
absentSince: absentSince,
|
|
}
|
|
}
|
|
default:
|
|
if len(r.PeerReconnectVerifiers) > maxRoomSize-1 {
|
|
validReservations = false
|
|
break
|
|
}
|
|
if len(r.PeerReconnectVerifiers) != 0 {
|
|
reservations = make(map[string]peerReservation, len(r.PeerReconnectVerifiers))
|
|
}
|
|
for peerID, encodedVerifier := range r.PeerReconnectVerifiers {
|
|
verifier, verifierOK := reconnectVerifierFromSnapshot(encodedVerifier)
|
|
if !validRelayID(peerID, maxPeerIDLength) ||
|
|
peerID == r.HostPeerID ||
|
|
!verifierOK {
|
|
validReservations = false
|
|
break
|
|
}
|
|
reservations[peerID] = peerReservation{
|
|
verifier: verifier,
|
|
absentSince: now,
|
|
}
|
|
}
|
|
}
|
|
if !validReservations {
|
|
skipped++
|
|
continue
|
|
}
|
|
|
|
restored[r.SessionID] = &Room{
|
|
SessionID: r.SessionID,
|
|
HostPeerID: r.HostPeerID,
|
|
ProtocolVersion: r.ProtocolVersion,
|
|
hostVerifier: hostVerifier,
|
|
peerReservations: reservations,
|
|
Peers: make(map[string]*Client),
|
|
CreatedAt: r.CreatedAt,
|
|
LastActivityAt: r.LastActivityAt,
|
|
}
|
|
}
|
|
s.mu.Lock()
|
|
s.rooms = restored
|
|
s.mu.Unlock()
|
|
log.Printf("snapshot: loaded %d rooms, skipped %d invalid or expired rooms", len(restored), skipped)
|
|
return rewriteReservations, nil
|
|
}
|
|
|
|
func (s *Server) cleanupLoop() {
|
|
ticker := time.NewTicker(cleanupInterval)
|
|
defer ticker.Stop()
|
|
for range ticker.C {
|
|
s.runCleanupStep(time.Now())
|
|
}
|
|
}
|
|
|
|
func (s *Server) runCleanupStep(now time.Time) {
|
|
s.mu.Lock()
|
|
var expiredClients []*Client
|
|
for id, room := range s.rooms {
|
|
room.mu.Lock()
|
|
roomChanged := pruneExpiredPeerReservationsLocked(room, now)
|
|
empty := len(room.Peers) == 0
|
|
age := now.Sub(room.CreatedAt)
|
|
idle := now.Sub(room.LastActivityAt)
|
|
expired := age > roomMaxAge
|
|
remove := (empty && idle > emptyRoomMaxAge) || expired
|
|
if remove {
|
|
room.closing = true
|
|
if expired && !empty {
|
|
for _, client := range room.Peers {
|
|
expiredClients = append(expiredClients, client)
|
|
}
|
|
clear(room.Peers)
|
|
}
|
|
log.Printf("cleanup: removing room %s (empty=%v, idle=%v, age=%v)", id, empty, idle, age)
|
|
s.removeRoomLocked(id, room)
|
|
roomChanged = true
|
|
}
|
|
if roomChanged {
|
|
s.snap.recordMutation()
|
|
}
|
|
room.mu.Unlock()
|
|
}
|
|
roomCount := len(s.rooms)
|
|
s.mu.Unlock()
|
|
|
|
for _, client := range expiredClients {
|
|
client.close()
|
|
}
|
|
if err := s.logs.cleanup(now); err != nil {
|
|
s.logRemovalError("logs", "cleanup", err)
|
|
}
|
|
if err := s.posters.cleanup(now); err != nil {
|
|
s.logRemovalError("posters", "cleanup", err)
|
|
}
|
|
s.posterUploads.cleanup(now)
|
|
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), roomCount)
|
|
s.conns.mu.Unlock()
|
|
}
|
|
|
|
func (s *Server) handlePostLogs(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
|
|
ip, err := s.clientIPs.resolve(r)
|
|
if err != nil {
|
|
http.Error(w, "Invalid client address", http.StatusBadRequest)
|
|
return
|
|
}
|
|
s.logs.mu.Lock()
|
|
if last, ok := s.logs.rateLimit[ip]; ok && time.Since(last) < logRateInterval {
|
|
s.logs.mu.Unlock()
|
|
http.Error(w, "Rate limited: 1 upload per minute", http.StatusTooManyRequests)
|
|
return
|
|
}
|
|
s.logs.rateLimit[ip] = time.Now()
|
|
s.logs.mu.Unlock()
|
|
|
|
body, err := io.ReadAll(io.LimitReader(r.Body, maxLogSize+1))
|
|
if err != nil {
|
|
http.Error(w, "Failed to read body", http.StatusBadRequest)
|
|
return
|
|
}
|
|
if len(body) > maxLogSize {
|
|
http.Error(w, "Log too large (max 1MB)", http.StatusRequestEntityTooLarge)
|
|
return
|
|
}
|
|
if len(body) == 0 {
|
|
http.Error(w, "Empty body", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
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
|
|
}
|
|
var removalErr *artifactRemovalError
|
|
if errors.As(err, &removalErr) {
|
|
s.logRemovalError("logs", "store", err)
|
|
} else {
|
|
log.Printf("logs: failed to store from %s: %v", ip, err)
|
|
}
|
|
http.Error(w, "Failed to store log", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
|
|
log.Printf("logs: stored %d bytes from %s", entry.Size, ip)
|
|
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]string{"id": id})
|
|
}
|
|
|
|
func (s *Server) handleGetLogs(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Cache-Control", "private, no-store")
|
|
if r.Method != http.MethodGet {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
|
|
type lookupResult struct {
|
|
entry artifactEntry
|
|
data []byte
|
|
status int
|
|
message string
|
|
}
|
|
lookup := func() lookupResult {
|
|
select {
|
|
case s.logLookups <- struct{}{}:
|
|
defer func() { <-s.logLookups }()
|
|
default:
|
|
return lookupResult{
|
|
status: http.StatusTooManyRequests,
|
|
message: "Too many concurrent lookups",
|
|
}
|
|
}
|
|
|
|
source, err := s.clientIPs.resolve(r)
|
|
if err != nil {
|
|
return lookupResult{
|
|
status: http.StatusBadRequest,
|
|
message: "Invalid client address",
|
|
}
|
|
}
|
|
id := strings.TrimPrefix(r.URL.Path, "/logs/")
|
|
entry, ok, err := s.logs.lookup(id, time.Now())
|
|
if err != nil {
|
|
s.logRemovalError("logs", "lookup", err)
|
|
return lookupResult{
|
|
status: http.StatusInternalServerError,
|
|
message: "Failed to retrieve log",
|
|
}
|
|
}
|
|
if !ok {
|
|
if !s.logs.allowFailedLookup(source, time.Now()) {
|
|
return lookupResult{
|
|
status: http.StatusTooManyRequests,
|
|
message: "Too many failed lookups",
|
|
}
|
|
}
|
|
return lookupResult{status: http.StatusNotFound, message: "Not found"}
|
|
}
|
|
|
|
data, err := os.ReadFile(s.logs.filePath(id))
|
|
if err != nil {
|
|
return lookupResult{status: http.StatusNotFound, message: "Not found"}
|
|
}
|
|
return lookupResult{entry: entry, data: data}
|
|
}()
|
|
|
|
controller := http.NewResponseController(w)
|
|
if err := controller.SetWriteDeadline(time.Now().Add(httpResponseWriteTimeout)); err != nil &&
|
|
!errors.Is(err, http.ErrNotSupported) {
|
|
log.Printf("logs: failed to set response write deadline")
|
|
}
|
|
|
|
if lookup.status != 0 {
|
|
http.Error(w, lookup.message, lookup.status)
|
|
return
|
|
}
|
|
|
|
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
|
w.Header().Set("Content-Length", strconv.FormatInt(lookup.entry.Size, 10))
|
|
if written, err := w.Write(lookup.data); err != nil || written != len(lookup.data) {
|
|
log.Printf("logs: response write failed")
|
|
}
|
|
}
|
|
|
|
var errPosterBodyReadTimeout = errors.New("poster body read timeout")
|
|
|
|
func readPosterBody(body io.ReadCloser, maxBytes int64, timeout time.Duration) ([]byte, error) {
|
|
timedOut := make(chan struct{})
|
|
timer := time.AfterFunc(timeout, func() {
|
|
close(timedOut)
|
|
_ = body.Close()
|
|
})
|
|
data, err := io.ReadAll(io.LimitReader(body, maxBytes+1))
|
|
if timer.Stop() {
|
|
return data, err
|
|
}
|
|
<-timedOut
|
|
return nil, errPosterBodyReadTimeout
|
|
}
|
|
|
|
func (s *Server) handlePostPosters(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodPost {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
|
|
ip, err := s.clientIPs.resolve(r)
|
|
if err != nil {
|
|
http.Error(w, "Invalid client address", http.StatusBadRequest)
|
|
return
|
|
}
|
|
if !s.posterUploads.tryStart(ip, time.Now()) {
|
|
http.Error(w, "Too many poster uploads", http.StatusTooManyRequests)
|
|
return
|
|
}
|
|
defer s.posterUploads.finish()
|
|
|
|
timeout := s.posterBodyReadTimeout
|
|
if timeout <= 0 {
|
|
timeout = posterUploadReadTimeout
|
|
}
|
|
readDeadline := http.NewResponseController(w)
|
|
if err := readDeadline.SetReadDeadline(time.Now().Add(timeout)); err == nil {
|
|
defer readDeadline.SetReadDeadline(time.Time{})
|
|
}
|
|
body, err := readPosterBody(r.Body, maxPosterSize, timeout)
|
|
var timeoutErr net.Error
|
|
if errors.Is(err, errPosterBodyReadTimeout) || errors.As(err, &timeoutErr) && timeoutErr.Timeout() {
|
|
http.Error(w, "Request body timeout", http.StatusRequestTimeout)
|
|
return
|
|
}
|
|
if err != nil {
|
|
http.Error(w, "Failed to read body", http.StatusBadRequest)
|
|
return
|
|
}
|
|
if len(body) > maxPosterSize {
|
|
http.Error(w, "Poster too large (max 5MB)", http.StatusRequestEntityTooLarge)
|
|
return
|
|
}
|
|
if len(body) == 0 {
|
|
http.Error(w, "Empty body", http.StatusBadRequest)
|
|
return
|
|
}
|
|
|
|
contentType := http.DetectContentType(body)
|
|
if _, ok := posterExtForContentType(contentType); !ok {
|
|
http.Error(w, "Unsupported media type", http.StatusUnsupportedMediaType)
|
|
return
|
|
}
|
|
|
|
id, entry, err := s.posters.store(body, contentType, time.Now())
|
|
if err != nil {
|
|
var removalErr *artifactRemovalError
|
|
if errors.As(err, &removalErr) {
|
|
s.logRemovalError("posters", "store", err)
|
|
} else {
|
|
log.Printf("posters: failed to store from %s: %v", ip, err)
|
|
}
|
|
http.Error(w, "Failed to store poster", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
|
|
url := "/posters/" + entry.Filename
|
|
log.Printf("posters: stored %s (%d bytes) from %s", id, entry.Size, ip)
|
|
|
|
w.Header().Set("Content-Type", "application/json")
|
|
json.NewEncoder(w).Encode(map[string]any{
|
|
"id": id,
|
|
"url": url,
|
|
"expiresIn": int(s.posters.maxAge.Seconds()),
|
|
})
|
|
}
|
|
|
|
func (s *Server) handleGetPosters(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodGet {
|
|
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
|
|
filename := strings.TrimPrefix(r.URL.Path, "/posters/")
|
|
entry, ok, err := s.posters.lookup(filename, time.Now())
|
|
if err != nil {
|
|
s.logRemovalError("posters", "lookup", err)
|
|
http.Error(w, "Failed to retrieve poster", http.StatusInternalServerError)
|
|
return
|
|
}
|
|
if !ok {
|
|
http.Error(w, "Not found", http.StatusNotFound)
|
|
return
|
|
}
|
|
|
|
f, err := os.Open(s.posters.filePath(entry.Filename))
|
|
if err != nil {
|
|
http.Error(w, "Not found", http.StatusNotFound)
|
|
return
|
|
}
|
|
defer f.Close()
|
|
|
|
remaining := int(time.Until(entry.ExpiresAt).Seconds())
|
|
if remaining < 0 {
|
|
remaining = 0
|
|
}
|
|
w.Header().Set("Cache-Control", "public, max-age="+strconv.Itoa(remaining))
|
|
w.Header().Set("Content-Type", entry.ContentType)
|
|
w.Header().Set("Content-Length", strconv.FormatInt(entry.Size, 10))
|
|
http.ServeContent(w, r, entry.Filename, entry.CreatedAt, f)
|
|
}
|
|
func (s *Server) handleWS(w http.ResponseWriter, r *http.Request) {
|
|
ip, err := s.clientIPs.resolve(r)
|
|
if err != nil {
|
|
http.Error(w, "Invalid client address", http.StatusBadRequest)
|
|
return
|
|
}
|
|
// Retained-room ownership uses the same canonical source key as connection admission.
|
|
quotaOwnerKey := ip
|
|
|
|
if !s.conns.tryConnect(ip) {
|
|
http.Error(w, "Too many connections", http.StatusTooManyRequests)
|
|
return
|
|
}
|
|
defer s.conns.disconnect(ip)
|
|
|
|
conn, err := upgrader.Upgrade(w, r, nil)
|
|
if err != nil {
|
|
log.Printf("upgrade error: %v", err)
|
|
return
|
|
}
|
|
defer conn.Close()
|
|
|
|
conn.SetReadLimit(maxMessageSize)
|
|
conn.SetReadDeadline(time.Now().Add(pongWait))
|
|
conn.SetPongHandler(func(string) error {
|
|
conn.SetReadDeadline(time.Now().Add(pongWait))
|
|
return nil
|
|
})
|
|
|
|
client := newClient(conn)
|
|
defer client.close()
|
|
|
|
rl := newRateLimiter(rateBurst, rateSustained)
|
|
var currentRoom *Room
|
|
var currentPeerID string
|
|
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 when this client is still authoritative. A
|
|
// displaced client's defer must neither remove the replacement nor start
|
|
// its reservation's absence clock.
|
|
defer func() {
|
|
if currentRoom != nil && currentPeerID != "" {
|
|
currentRoom.mu.Lock()
|
|
closing := currentRoom.closing
|
|
stale := currentRoom.Peers[currentPeerID] != client
|
|
if !closing && !stale {
|
|
now := time.Now()
|
|
delete(currentRoom.Peers, currentPeerID)
|
|
if currentRoom.ProtocolVersion == relayProtocolVersion &&
|
|
currentPeerID != currentRoom.HostPeerID {
|
|
if reservation, ok := currentRoom.peerReservations[currentPeerID]; ok {
|
|
reservation.absentSince = now
|
|
currentRoom.peerReservations[currentPeerID] = reservation
|
|
}
|
|
}
|
|
currentRoom.LastActivityAt = now
|
|
s.snap.recordMutation()
|
|
}
|
|
currentRoom.mu.Unlock()
|
|
if !closing && !stale {
|
|
currentRoom.broadcastExcept(currentPeerID, serverMsg{
|
|
Type: relayTypePeerLeft,
|
|
PeerID: currentPeerID,
|
|
})
|
|
}
|
|
log.Printf("peer %s left room %s (closing=%v, stale=%v)", currentPeerID, currentRoom.SessionID, closing, stale)
|
|
}
|
|
}()
|
|
|
|
for {
|
|
_, raw, err := conn.ReadMessage()
|
|
if err != nil {
|
|
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway, websocket.CloseNormalClosure) {
|
|
log.Printf("read error: %v", err)
|
|
}
|
|
return
|
|
}
|
|
|
|
if !rl.allow() {
|
|
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: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid JSON"})
|
|
continue
|
|
}
|
|
|
|
switch msg.Type {
|
|
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 msg.ProtocolVersion != legacyRelayProtocolVersion && msg.ProtocolVersion != relayProtocolVersion {
|
|
client.sendJSON(serverMsg{
|
|
Type: relayTypeError,
|
|
Code: relayErrorProtocolMismatch,
|
|
Message: "Unsupported relay protocol version",
|
|
ProtocolVersion: relayProtocolVersion,
|
|
})
|
|
continue
|
|
}
|
|
if rejectRoomTransition() {
|
|
continue
|
|
}
|
|
|
|
reconnectToken := msg.ReconnectToken
|
|
var hostVerifier reconnectVerifier
|
|
if reconnectToken == "" {
|
|
if msg.ProtocolVersion != legacyRelayProtocolVersion {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Modern room creation requires a reconnect token"})
|
|
continue
|
|
}
|
|
var err error
|
|
reconnectToken, hostVerifier, err = mintReconnectToken()
|
|
if err != nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to create room"})
|
|
continue
|
|
}
|
|
} else {
|
|
var ok bool
|
|
hostVerifier, ok = reconnectVerifierFromToken(reconnectToken)
|
|
if !ok {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid reconnect token"})
|
|
continue
|
|
}
|
|
}
|
|
|
|
var rejection *serverMsg
|
|
var oldHostClient *Client
|
|
var hostWasAbsent bool
|
|
s.mu.Lock()
|
|
existing := s.rooms[msg.SessionID]
|
|
if existing != nil {
|
|
existing.mu.Lock()
|
|
idempotentModernCreate :=
|
|
!existing.closing &&
|
|
msg.ProtocolVersion == relayProtocolVersion &&
|
|
existing.ProtocolVersion == relayProtocolVersion &&
|
|
msg.PeerID == existing.HostPeerID &&
|
|
reconnectVerifierMatches(existing.hostVerifier, hostVerifier)
|
|
if idempotentModernCreate {
|
|
oldHostClient = existing.Peers[msg.PeerID]
|
|
hostWasAbsent = oldHostClient == nil
|
|
existing.Peers[msg.PeerID] = client
|
|
existing.LastActivityAt = time.Now()
|
|
peers := existing.peerIDs()
|
|
s.snap.recordMutation()
|
|
existing.mu.Unlock()
|
|
s.mu.Unlock()
|
|
|
|
currentRoom = existing
|
|
currentPeerID = msg.PeerID
|
|
if oldHostClient != nil && oldHostClient != client {
|
|
oldHostClient.close()
|
|
}
|
|
existingPeers := make([]string, 0, len(peers)-1)
|
|
for _, peerID := range peers {
|
|
if peerID != msg.PeerID {
|
|
existingPeers = append(existingPeers, peerID)
|
|
}
|
|
}
|
|
client.sendJSON(serverMsg{
|
|
Type: relayTypeCreated,
|
|
SessionID: msg.SessionID,
|
|
HostPeerID: msg.PeerID,
|
|
ReconnectToken: reconnectToken,
|
|
ProtocolVersion: relayProtocolVersion,
|
|
Peers: existingPeers,
|
|
})
|
|
if hostWasAbsent {
|
|
existing.broadcastExcept(msg.PeerID, serverMsg{
|
|
Type: relayTypePeerJoined,
|
|
PeerID: msg.PeerID,
|
|
})
|
|
}
|
|
continue
|
|
}
|
|
authorizedLegacyReplacement :=
|
|
len(existing.Peers) == 0 &&
|
|
!existing.closing &&
|
|
msg.ProtocolVersion == legacyRelayProtocolVersion &&
|
|
existing.ProtocolVersion == legacyRelayProtocolVersion &&
|
|
msg.PeerID == existing.HostPeerID &&
|
|
msg.ReconnectToken != "" &&
|
|
reconnectVerifierMatches(existing.hostVerifier, hostVerifier)
|
|
existing.mu.Unlock()
|
|
if !authorizedLegacyReplacement {
|
|
rejection = &serverMsg{Type: relayTypeError, Code: relayErrorRoomExists, Message: "Room already exists"}
|
|
}
|
|
} else if len(s.rooms) >= maxRetainedRooms {
|
|
rejection = &serverMsg{Type: relayTypeError, Code: relayErrorRateLimited, Message: "Too many retained rooms"}
|
|
}
|
|
if rejection == nil {
|
|
var reserved bool
|
|
if existing == nil {
|
|
reserved = s.conns.tryCreateRoom(quotaOwnerKey)
|
|
} else {
|
|
reserved = s.conns.tryCreateRoomReplacing(quotaOwnerKey, existing.quotaOwnerKey)
|
|
}
|
|
if !reserved {
|
|
rejection = &serverMsg{Type: relayTypeError, Code: relayErrorRateLimited, Message: "Too many rooms created"}
|
|
}
|
|
}
|
|
if rejection != nil {
|
|
s.mu.Unlock()
|
|
client.sendJSON(*rejection)
|
|
continue
|
|
}
|
|
if existing != nil {
|
|
s.removeRoomLocked(msg.SessionID, existing)
|
|
}
|
|
now := time.Now()
|
|
room := &Room{
|
|
SessionID: msg.SessionID,
|
|
HostPeerID: msg.PeerID,
|
|
ProtocolVersion: msg.ProtocolVersion,
|
|
hostVerifier: hostVerifier,
|
|
Peers: map[string]*Client{msg.PeerID: client},
|
|
quotaOwnerKey: quotaOwnerKey,
|
|
CreatedAt: now,
|
|
LastActivityAt: now,
|
|
}
|
|
|
|
s.rooms[msg.SessionID] = room
|
|
s.snap.recordMutation()
|
|
s.mu.Unlock()
|
|
|
|
currentRoom = room
|
|
currentPeerID = msg.PeerID
|
|
log.Printf("room %s created by %s", msg.SessionID, msg.PeerID)
|
|
client.sendJSON(serverMsg{
|
|
Type: relayTypeCreated,
|
|
SessionID: msg.SessionID,
|
|
HostPeerID: msg.PeerID,
|
|
ReconnectToken: reconnectToken,
|
|
ProtocolVersion: msg.ProtocolVersion,
|
|
})
|
|
|
|
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 msg.ProtocolVersion != legacyRelayProtocolVersion && msg.ProtocolVersion != relayProtocolVersion {
|
|
client.sendJSON(serverMsg{
|
|
Type: relayTypeError,
|
|
Code: relayErrorProtocolMismatch,
|
|
Message: "Unsupported relay protocol version",
|
|
ProtocolVersion: relayProtocolVersion,
|
|
})
|
|
continue
|
|
}
|
|
if rejectRoomTransition() {
|
|
continue
|
|
}
|
|
|
|
newToken, newVerifier, err := mintReconnectToken()
|
|
if err != nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to join room"})
|
|
continue
|
|
}
|
|
presentedVerifier, tokenValid := reconnectVerifierFromToken(msg.ReconnectToken)
|
|
|
|
s.mu.RLock()
|
|
room, exists := s.rooms[msg.SessionID]
|
|
if !exists {
|
|
s.mu.RUnlock()
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomNotFound, Message: "Room does not exist"})
|
|
continue
|
|
}
|
|
if s.beforeJoinRoomLock != nil {
|
|
s.beforeJoinRoomLock()
|
|
}
|
|
room.mu.Lock()
|
|
if s.rooms[msg.SessionID] != room {
|
|
room.mu.Unlock()
|
|
s.mu.RUnlock()
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomNotFound, Message: "Room does not exist"})
|
|
continue
|
|
}
|
|
s.mu.RUnlock()
|
|
if room.closing {
|
|
room.mu.Unlock()
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomNotFound, Message: "Room does not exist"})
|
|
continue
|
|
}
|
|
if room.ProtocolVersion != msg.ProtocolVersion {
|
|
room.mu.Unlock()
|
|
client.sendJSON(serverMsg{
|
|
Type: relayTypeError,
|
|
Code: relayErrorProtocolMismatch,
|
|
Message: "Client and room protocol versions are incompatible",
|
|
ProtocolVersion: room.ProtocolVersion,
|
|
})
|
|
continue
|
|
}
|
|
|
|
if room.ProtocolVersion == relayProtocolVersion &&
|
|
pruneExpiredPeerReservationsLocked(room, time.Now()) {
|
|
s.snap.recordMutation()
|
|
}
|
|
existingClient, occupied := room.Peers[msg.PeerID]
|
|
reservation, identityReserved := room.peerReservations[msg.PeerID]
|
|
responseToken := newToken
|
|
responseVerifier := newVerifier
|
|
authorized := false
|
|
if room.ProtocolVersion == relayProtocolVersion {
|
|
switch {
|
|
case msg.PeerID == room.HostPeerID:
|
|
authorized = tokenValid && reconnectVerifierMatches(room.hostVerifier, presentedVerifier)
|
|
responseToken = msg.ReconnectToken
|
|
responseVerifier = room.hostVerifier
|
|
case identityReserved:
|
|
authorized = !reservation.releasePending &&
|
|
tokenValid &&
|
|
reconnectVerifierMatches(reservation.verifier, presentedVerifier)
|
|
responseToken = msg.ReconnectToken
|
|
responseVerifier = reservation.verifier
|
|
case occupied:
|
|
authorized = false
|
|
default:
|
|
authorized = tokenValid
|
|
responseToken = msg.ReconnectToken
|
|
responseVerifier = presentedVerifier
|
|
}
|
|
} else if msg.PeerID == room.HostPeerID {
|
|
switch {
|
|
case tokenValid && reconnectVerifierMatches(room.hostVerifier, presentedVerifier):
|
|
authorized = true
|
|
responseToken = msg.ReconnectToken
|
|
responseVerifier = room.hostVerifier
|
|
case msg.ReconnectToken == "" &&
|
|
!occupied &&
|
|
room.quotaOwnerKey != "" &&
|
|
room.quotaOwnerKey == quotaOwnerKey:
|
|
// Tokenless host reconnect is retained only for unversioned rooms,
|
|
// only within this process, and only from the creating source.
|
|
authorized = true
|
|
responseToken = ""
|
|
responseVerifier = room.hostVerifier
|
|
}
|
|
} else {
|
|
// Legacy guests have no durable proof. Never let one replace a live
|
|
// identity; disconnected identity reuse remains confined to legacy rooms.
|
|
authorized = !occupied
|
|
}
|
|
|
|
if !authorized {
|
|
room.mu.Unlock()
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorPeerIdUnavailable, Message: "Peer ID is unavailable"})
|
|
continue
|
|
}
|
|
if !occupied {
|
|
roomFull := false
|
|
if room.ProtocolVersion == relayProtocolVersion {
|
|
if msg.PeerID != room.HostPeerID && !identityReserved {
|
|
roomFull = len(room.peerReservations) >= maxRoomSize-1
|
|
}
|
|
} else {
|
|
admissionLimit := maxRoomSize
|
|
_, hostConnected := room.Peers[room.HostPeerID]
|
|
if msg.PeerID != room.HostPeerID && !hostConnected {
|
|
admissionLimit--
|
|
}
|
|
roomFull = len(room.Peers) >= admissionLimit
|
|
}
|
|
if roomFull {
|
|
room.mu.Unlock()
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorRoomFull, Message: "Room is full"})
|
|
continue
|
|
}
|
|
}
|
|
|
|
room.Peers[msg.PeerID] = client
|
|
if room.ProtocolVersion == relayProtocolVersion && msg.PeerID != room.HostPeerID {
|
|
if identityReserved {
|
|
reservation.absentSince = time.Time{}
|
|
reservation.releasePending = false
|
|
reservation.releaseClient = nil
|
|
reservation.verifier = responseVerifier
|
|
} else {
|
|
reservation = peerReservation{verifier: responseVerifier}
|
|
}
|
|
if room.peerReservations == nil {
|
|
room.peerReservations = make(map[string]peerReservation)
|
|
}
|
|
room.peerReservations[msg.PeerID] = reservation
|
|
}
|
|
room.LastActivityAt = time.Now()
|
|
peers := room.peerIDs()
|
|
hostPeerID := room.HostPeerID
|
|
roomProtocolVersion := room.ProtocolVersion
|
|
s.snap.recordMutation()
|
|
room.mu.Unlock()
|
|
|
|
currentRoom = room
|
|
currentPeerID = msg.PeerID
|
|
if occupied && existingClient != client {
|
|
existingClient.close()
|
|
}
|
|
log.Printf("peer %s joined room %s", msg.PeerID, msg.SessionID)
|
|
|
|
existingPeers := make([]string, 0, len(peers)-1)
|
|
for _, peerID := range peers {
|
|
if peerID != msg.PeerID {
|
|
existingPeers = append(existingPeers, peerID)
|
|
}
|
|
}
|
|
client.sendJSON(serverMsg{
|
|
Type: relayTypeJoined,
|
|
SessionID: msg.SessionID,
|
|
HostPeerID: hostPeerID,
|
|
ReconnectToken: responseToken,
|
|
ProtocolVersion: roomProtocolVersion,
|
|
Peers: existingPeers,
|
|
})
|
|
room.broadcastExcept(msg.PeerID, serverMsg{Type: relayTypePeerJoined, PeerID: msg.PeerID})
|
|
|
|
case relayTypeLeave:
|
|
if currentRoom == nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"})
|
|
continue
|
|
}
|
|
room := currentRoom
|
|
presentedVerifier, tokenValid := reconnectVerifierFromToken(msg.ReconnectToken)
|
|
if s.beforeLeaveRoomLock != nil {
|
|
s.beforeLeaveRoomLock()
|
|
}
|
|
room.mu.Lock()
|
|
currentClient := room.Peers[currentPeerID] == client
|
|
isGuest := currentPeerID != room.HostPeerID
|
|
authorized := currentClient && isGuest && !room.closing && msg.ProtocolVersion == room.ProtocolVersion
|
|
reservation, reservationOK := room.peerReservations[currentPeerID]
|
|
if authorized && room.ProtocolVersion == relayProtocolVersion {
|
|
authorized = reservationOK &&
|
|
!reservation.releasePending &&
|
|
tokenValid &&
|
|
reconnectVerifierMatches(reservation.verifier, presentedVerifier)
|
|
}
|
|
if !authorized {
|
|
room.mu.Unlock()
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorPeerIdUnavailable, Message: "Unable to release peer identity"})
|
|
continue
|
|
}
|
|
|
|
releasedPeerID := currentPeerID
|
|
if room.ProtocolVersion == legacyRelayProtocolVersion {
|
|
delete(room.Peers, releasedPeerID)
|
|
room.LastActivityAt = time.Now()
|
|
s.snap.recordMutation()
|
|
room.mu.Unlock()
|
|
currentRoom = nil
|
|
currentPeerID = ""
|
|
} else {
|
|
previousActivity := room.LastActivityAt
|
|
leaveActivity := time.Now()
|
|
reservation.releasePending = true
|
|
reservation.releaseClient = client
|
|
room.peerReservations[releasedPeerID] = reservation
|
|
room.LastActivityAt = leaveActivity
|
|
ticket := s.snap.recordTerminalMutation(func(persistErr error) terminalMutationOutcome {
|
|
s.mu.RLock()
|
|
authoritativeRoom := s.rooms[room.SessionID] == room
|
|
room.mu.Lock()
|
|
currentReservation, stillReserved := room.peerReservations[releasedPeerID]
|
|
ownsRelease := authoritativeRoom &&
|
|
!room.closing &&
|
|
stillReserved &&
|
|
currentReservation.releasePending &&
|
|
currentReservation.releaseClient == client &&
|
|
room.Peers[releasedPeerID] == client
|
|
if !ownsRelease {
|
|
room.mu.Unlock()
|
|
s.mu.RUnlock()
|
|
return terminalMutationOutcome{err: persistErr, deliver: false}
|
|
}
|
|
if persistErr == nil {
|
|
delete(room.Peers, releasedPeerID)
|
|
delete(room.peerReservations, releasedPeerID)
|
|
} else {
|
|
currentReservation.releasePending = false
|
|
currentReservation.releaseClient = nil
|
|
room.peerReservations[releasedPeerID] = currentReservation
|
|
if room.LastActivityAt.Equal(leaveActivity) {
|
|
room.LastActivityAt = previousActivity
|
|
}
|
|
// The failed attempt captured the pending omission. Record
|
|
// the restored reservation before a queued later capture.
|
|
s.snap.recordMutation()
|
|
}
|
|
room.mu.Unlock()
|
|
s.mu.RUnlock()
|
|
return terminalMutationOutcome{err: persistErr, deliver: true}
|
|
})
|
|
room.mu.Unlock()
|
|
|
|
outcome := s.snap.waitForDurable(ticket)
|
|
if !outcome.deliver {
|
|
continue
|
|
}
|
|
if outcome.err != nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to persist released peer identity"})
|
|
continue
|
|
}
|
|
currentRoom = nil
|
|
currentPeerID = ""
|
|
}
|
|
|
|
if s.beforeTerminalDelivery != nil {
|
|
s.beforeTerminalDelivery()
|
|
}
|
|
client.sendJSON(serverMsg{
|
|
Type: relayTypeLeft,
|
|
SessionID: room.SessionID,
|
|
PeerID: releasedPeerID,
|
|
ProtocolVersion: room.ProtocolVersion,
|
|
})
|
|
room.broadcastExcept(releasedPeerID, serverMsg{Type: relayTypePeerLeft, PeerID: releasedPeerID})
|
|
|
|
case relayTypeEndSession:
|
|
if currentRoom == nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"})
|
|
continue
|
|
}
|
|
room := currentRoom
|
|
presentedVerifier, tokenValid := reconnectVerifierFromToken(msg.ReconnectToken)
|
|
s.mu.Lock()
|
|
room.mu.Lock()
|
|
authorized :=
|
|
s.rooms[room.SessionID] == room &&
|
|
!room.closing &&
|
|
currentPeerID == room.HostPeerID &&
|
|
room.Peers[currentPeerID] == client &&
|
|
msg.ProtocolVersion == room.ProtocolVersion
|
|
if authorized && room.ProtocolVersion == relayProtocolVersion {
|
|
authorized = tokenValid && reconnectVerifierMatches(room.hostVerifier, presentedVerifier)
|
|
}
|
|
if !authorized {
|
|
room.mu.Unlock()
|
|
s.mu.Unlock()
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorPeerIdUnavailable, Message: "Unable to end room"})
|
|
continue
|
|
}
|
|
room.closing = true
|
|
guests := make([]*Client, 0, len(room.Peers)-1)
|
|
for peerID, peerClient := range room.Peers {
|
|
if peerID != currentPeerID {
|
|
guests = append(guests, peerClient)
|
|
}
|
|
}
|
|
s.removeRoomLocked(room.SessionID, room)
|
|
ticket := s.snap.recordTerminalMutation(nil)
|
|
room.mu.Unlock()
|
|
s.mu.Unlock()
|
|
|
|
// A successful outcome means the file-synced atomic rename committed.
|
|
// Supported filesystems also complete parent-directory sync before
|
|
// this barrier; post-rename sync degradation is warning-only.
|
|
outcome := s.snap.waitForDurable(ticket)
|
|
if outcome.err != nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unable to persist ended room"})
|
|
room.mu.Lock()
|
|
clear(room.Peers)
|
|
clear(room.peerReservations)
|
|
room.mu.Unlock()
|
|
currentRoom = nil
|
|
currentPeerID = ""
|
|
for _, guest := range guests {
|
|
guest.close()
|
|
}
|
|
continue
|
|
}
|
|
if s.beforeTerminalDelivery != nil {
|
|
s.beforeTerminalDelivery()
|
|
}
|
|
endedMessage := serverMsg{
|
|
Type: relayTypeEnded,
|
|
SessionID: room.SessionID,
|
|
ProtocolVersion: room.ProtocolVersion,
|
|
}
|
|
client.sendJSON(endedMessage)
|
|
for _, guest := range guests {
|
|
guest.sendJSONAndWait(endedMessage)
|
|
}
|
|
|
|
room.mu.Lock()
|
|
clear(room.Peers)
|
|
clear(room.peerReservations)
|
|
room.mu.Unlock()
|
|
currentRoom = nil
|
|
currentPeerID = ""
|
|
for _, guest := range guests {
|
|
guest.close()
|
|
}
|
|
|
|
case relayTypeBroadcast:
|
|
if currentRoom == nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"})
|
|
continue
|
|
}
|
|
if !currentRoom.broadcastFrom(currentPeerID, client, serverMsg{
|
|
Type: relayTypeMessage,
|
|
From: currentPeerID,
|
|
Payload: msg.Payload,
|
|
}) {
|
|
client.close()
|
|
return
|
|
}
|
|
|
|
case relayTypeSendTo:
|
|
if currentRoom == nil {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Not in a room"})
|
|
continue
|
|
}
|
|
if !validRelayID(msg.To, maxPeerIDLength) {
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Invalid to field"})
|
|
continue
|
|
}
|
|
switch currentRoom.sendFrom(currentPeerID, client, msg.To, serverMsg{
|
|
Type: relayTypeMessage,
|
|
From: currentPeerID,
|
|
Payload: msg.Payload,
|
|
}) {
|
|
case directedSenderUnavailable:
|
|
client.close()
|
|
return
|
|
case directedTargetMissing:
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorNotInRoom, Message: "Target peer not found"})
|
|
}
|
|
|
|
case relayTypePing:
|
|
client.sendJSON(serverMsg{Type: relayTypePong})
|
|
|
|
default:
|
|
client.sendJSON(serverMsg{Type: relayTypeError, Code: relayErrorInvalidMessage, Message: "Unknown message type"})
|
|
}
|
|
}
|
|
}
|
|
|
|
func newHTTPServer(addr string, handler http.Handler) *http.Server {
|
|
return &http.Server{
|
|
Addr: addr,
|
|
Handler: handler,
|
|
ReadTimeout: posterUploadReadTimeout,
|
|
WriteTimeout: httpResponseWriteTimeout,
|
|
MaxHeaderBytes: maxHTTPHeaderBytes,
|
|
}
|
|
}
|
|
|
|
func main() {
|
|
addr := flag.String("addr", ":8080", "Listen address")
|
|
logDir := flag.String("log-dir", "/data/logs", "Directory for log file storage")
|
|
posterDir := flag.String("poster-dir", "/data/posters", "Directory for Discord poster storage")
|
|
stateFile := flag.String("state-file", "/data/rooms.json", "Path to room snapshot file")
|
|
flag.Parse()
|
|
|
|
trustedProxyCIDRs, err := parseTrustedProxyCIDRs(os.Getenv("TRUSTED_PROXY_CIDRS"))
|
|
if err != nil {
|
|
log.Fatalf("invalid TRUSTED_PROXY_CIDRS")
|
|
}
|
|
clientIPs := newClientIPResolver(trustedProxyCIDRs)
|
|
srv := newServer(*logDir, *stateFile, *posterDir, clientIPs)
|
|
|
|
mux := http.NewServeMux()
|
|
mux.HandleFunc("/relay", srv.handleWS)
|
|
mux.HandleFunc("/health", func(w http.ResponseWriter, r *http.Request) {
|
|
w.WriteHeader(http.StatusOK)
|
|
w.Write([]byte("ok"))
|
|
})
|
|
mux.HandleFunc("/logs", srv.handlePostLogs)
|
|
mux.HandleFunc("/logs/", srv.handleGetLogs)
|
|
mux.HandleFunc("/posters", srv.handlePostPosters)
|
|
mux.HandleFunc("/posters/", srv.handleGetPosters)
|
|
registerOAuthRoutes(mux, srv.oauth)
|
|
|
|
httpSrv := newHTTPServer(*addr, mux)
|
|
|
|
serveErr := make(chan error, 1)
|
|
go func() {
|
|
log.Printf("Starting relay server on %s", *addr)
|
|
if err := httpSrv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
|
serveErr <- err
|
|
}
|
|
close(serveErr)
|
|
}()
|
|
|
|
sig := make(chan os.Signal, 1)
|
|
signal.Notify(sig, syscall.SIGINT, syscall.SIGTERM)
|
|
|
|
select {
|
|
case err, ok := <-serveErr:
|
|
if ok {
|
|
log.Fatalf("listen: %v", err)
|
|
}
|
|
case s := <-sig:
|
|
log.Printf("shutdown signal received (%s), draining...", s)
|
|
}
|
|
|
|
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
|
defer cancel()
|
|
if err := httpSrv.Shutdown(ctx); err != nil {
|
|
log.Printf("http shutdown: %v", err)
|
|
}
|
|
if err := srv.snap.flushAndStop(snapshotFlushTimeout); err != nil {
|
|
log.Printf("snapshot flush: %v", err)
|
|
}
|
|
log.Printf("shutdown complete")
|
|
}
|