feat(playback): stream kill switch + async over-cap enforcer

Add the enforcement layer on top of server-observed monitoring: a revocation
kill switch that stops any stream within ~120s and keeps it dead, plus an async
over-cap enforcer that drives kills off the live monitoring picture — entirely
off the per-segment hot path and with no client-protocol change.

- internal/streamrevoke: the central kill list. IsRevoked is a pure in-memory
  lookup safe on the request hot path; a Redis pub/sub + poll mirror keeps edge
  caches current, and a Postgres durable mirror lets kills survive a server
  restart AND a Redis flush so a restart-resilient stream cannot be reconstructed
  and re-served after being killed. A user revocation is a cutoff (kills tokens
  minted before it, spares post-reauth tokens), not a 24h ban.
- internal/streamenforcer: async over-cap brain — reads the monitoring snapshot
  and per-user limits, selects victims, and collapses every reason (exceeded
  limit, admin terminate, abuse) to the same action: write a revocation.
- Edge + native + jellycompat enforcement: proxy refuses revoked sessions on
  every request and cuts long direct-play/remux pours mid-stream; the transcode
  node guards both serve and the reconstruct path so a killed session is never
  re-spawned after a node restart; jellycompat serve surfaces close their
  kill-switch coverage holes.
- streamtoken.IssuedTime exposes the token iat the user-kill cutoff compares
  against; token IssuedTime + revocation guards wire through router, downloads,
  and admin terminate-by-id (with admin-list dedupe).
- Restore sendfile zero-copy on direct-play/remux byte counting so the monitor's
  served-byte accounting does not cost the sendfile fast path.
- migrations/sql: stream_revocations durable table.

Part of the stream monitoring & kill-switch epic.
This commit is contained in:
CoffeeKnyte
2026-07-29 12:26:03 +00:00
parent 22ffaad911
commit 0217cce2df
23 changed files with 2272 additions and 33 deletions
+63
View File
@@ -95,6 +95,9 @@ import (
"github.com/Silo-Server/silo-server/internal/secret"
"github.com/Silo-Server/silo-server/internal/sections"
"github.com/Silo-Server/silo-server/internal/server"
"github.com/Silo-Server/silo-server/internal/streamenforcer"
"github.com/Silo-Server/silo-server/internal/streammonitor"
"github.com/Silo-Server/silo-server/internal/streamrevoke"
"github.com/Silo-Server/silo-server/internal/subtitles"
"github.com/Silo-Server/silo-server/internal/taskmanager"
taskrepository "github.com/Silo-Server/silo-server/internal/taskmanager/repository"
@@ -668,9 +671,16 @@ func main() {
tracker.Cleanup(cleanupCtx)
}()
// Kill switch: edges consult this per request. Central writes revocations
// (over-cap, admin terminate, account revocation) which propagate here via
// Redis pub/sub + poll. Shares the offload Redis.
revStore := streamrevoke.New(streamrevoke.Options{Redis: redisClient, Bus: eventBus})
revStore.StartSync(appCtx)
var handler http.Handler
if mode == "proxy" {
srv := proxy.NewServer(watcher, tracker)
srv.SetRevocationStore(revStore)
handler = srv.Handler()
} else {
srv := transcodenode.NewServer(watcher, tracker)
@@ -682,6 +692,7 @@ func main() {
// Reclaim orphaned transcode dirs at boot and hourly thereafter, bound
// to appCtx so it stops on shutdown.
srv.StartOrphanSweeper(appCtx)
srv.SetRevocationStore(revStore)
handler = srv.Handler()
}
@@ -745,6 +756,20 @@ func main() {
defer func() { _ = apiRedisClient.Close() }()
}
// Central-side kill switch: written by the async enforcer, admin terminate,
// and account revocation; edges enforce it via Redis pub/sub + poll. Memory-
// only when Redis is absent (single-node integrated). The Postgres durable
// mirror lets the kill list survive a server restart AND a Redis flush, so a
// restart-resilient stream (PR #174) cannot be reconstructed and re-served
// after being killed. NewPostgresDurableStore returns a true nil interface
// when the pool is nil, degrading cleanly to memory/Redis-only.
streamRevocation := streamrevoke.New(streamrevoke.Options{
Redis: apiRedisClient,
Bus: eventBus,
Durable: streamrevoke.NewPostgresDurableStore(pool),
})
streamRevocation.StartSync(appCtx)
// Assigned below once the trusted-proxy config is seeded; captured by the
// OnServerSettingUpdated closure, which only runs on admin requests after
// startup completes.
@@ -765,6 +790,7 @@ func main() {
SecretCipher: dataCipher,
EventBus: eventBus,
RedisClient: apiRedisClient,
RevocationStore: streamRevocation,
LogStreamHub: logStreamHub,
RealtimeHub: realtimeHub,
EventsHub: eventsHub,
@@ -1702,6 +1728,36 @@ func main() {
}
deps.SessionSyncer = reconciler
// Async brain: off the hot path, trims over-cap streams by writing
// revocations the edges enforce. Source is the authoritative monitoring
// picture — Redis for multi-node, the in-process session manager for
// integrated. Fails open on limit-lookup errors.
enforcerUserRepo := auth.NewUserRepository(deps.DB)
// Union of (a) the in-process session manager — authoritative for streams
// this node serves directly, integrated or not — and (b) the edge Redis
// records — authoritative for offloaded/edge-served streams. Deduped by
// session id. Using the union (not "Redis when present, else manager")
// fixes the case where integrated mode has Redis configured but writes no
// silo:sessions:* keys, which would otherwise blind the enforcer.
// The mapping (Session → monitoring record, carrying route + client identity
// + progress) is shared with the admin session list via handlers so there is
// one Session→record definition. NodeName marks the serving host so
// integrated streams aren't shown node-less.
localSource := streammonitor.NewFuncSource(func(ctx context.Context) ([]nodesessions.SessionInfo, error) {
return handlers.LiveLocalSessions(sessionMgr, nodeIdentity), nil
})
var monitorSource streammonitor.Source = localSource
if apiRedisClient != nil {
monitorSource = streammonitor.NewMultiSource(localSource, streammonitor.NewRedisSource(apiRedisClient))
}
streamenforcer.New(monitorSource, func(ctx context.Context, userID int) (int, error) {
u, err := enforcerUserRepo.GetByID(ctx, userID)
if err != nil {
return 0, err
}
return u.MaxStreams, nil
}, streamRevocation, 0).Start(appCtx)
nodeURL := fmt.Sprintf("http://%s%s", nodeIdentity, cfg.Server.Listen)
heartbeatWriter = worker.NewHeartbeatWriter(deps.DB, nodeIdentity, mode, nodeURL)
}
@@ -2340,6 +2396,12 @@ func main() {
if compatServer != nil {
compatServer.SessionStore().DeleteByUserID(userID)
}
// Also revoke this user's live stream credentials so a session revocation
// (disable, password change, permission change) kills in-flight playback
// at the edge within one propagation/poll interval, not just compat login.
if err := streamRevocation.RevokeUser(ctx, userID, "user_sessions_revoked"); err != nil {
slog.Warn("revoke user streams failed", "user_id", userID, "error", err)
}
}
distFS, fsErr := fs.Sub(siloweb.DistFS, "dist")
@@ -2512,6 +2574,7 @@ func main() {
compatDeps.DetailSvc = detailSvc
compatDeps.FolderRepo = folderRepo
compatDeps.SessionMgr = sessionMgr
compatDeps.RevocationStore = streamRevocation
compatDeps.UserStoreProvider = userStoreProvider
compatDeps.WatchCompletionObserver = deps.WatchCompletionObserver
compatDeps.SettingsRepo = settingsRepo
@@ -5,6 +5,7 @@ import (
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"time"
@@ -12,6 +13,7 @@ import (
"github.com/google/uuid"
"github.com/Silo-Server/silo-server/internal/playback"
"github.com/Silo-Server/silo-server/internal/streamrevoke"
)
const (
@@ -20,7 +22,17 @@ const (
)
type AdminPlaybackControlHandler struct {
playback *PlaybackHandler
playback *PlaybackHandler
revocation *streamrevoke.Store
}
// SetRevocationStore wires the kill switch so an admin terminate revokes the
// stream credential (stops it at the edge and refuses reconnects), not just
// dispatches a cooperative realtime command. Optional; nil keeps prior behavior.
func (h *AdminPlaybackControlHandler) SetRevocationStore(store *streamrevoke.Store) {
if h != nil {
h.revocation = store
}
}
type playbackControlRequest struct {
@@ -145,9 +157,38 @@ func (h *AdminPlaybackControlHandler) handleSessionCommand(w http.ResponseWriter
return
}
var req playbackControlRequest
if err := decodeOptionalJSONBody(r, &req); err != nil {
writeError(w, http.StatusBadRequest, "bad_request", "Invalid request body")
return
}
// A terminate must stick even if the client ignores the realtime command:
// revoke the stream credential so the edge refuses further segments and any
// reconnect within one propagation/poll interval. Stop stays cooperative.
// The revocation is written BEFORE the local session lookup: the kill needs
// only the id, and the stream may exist only as an edge Redis record — e.g.
// after a central restart, or when a client that withholds progress let the
// in-memory session be reaped — which is precisely the stream an operator
// most needs to kill. Revoking an unknown id is harmless (idempotent, TTL'd).
if name == playback.CommandTerminate && h.revocation != nil {
if err := h.revocation.RevokeSession(r.Context(), sessionID, "admin_terminate"); err != nil {
slog.Warn("admin terminate: revoke stream failed", "session_id", sessionID, "error", err)
}
}
session, err := h.playback.sessionMgr.GetSession(sessionID)
if err != nil {
if errors.Is(err, playback.ErrSessionNotFound) {
// The revocation above still cut the stream at every serve surface;
// only the cooperative realtime command has no local session to go to.
if name == playback.CommandTerminate && h.revocation != nil {
writeJSON(w, http.StatusAccepted, playbackControlResponse{
CommandID: uuid.NewString(),
Status: "revoked",
})
return
}
writeError(w, http.StatusNotFound, "not_found", "Playback session not found")
return
}
@@ -155,12 +196,6 @@ func (h *AdminPlaybackControlHandler) handleSessionCommand(w http.ResponseWriter
return
}
var req playbackControlRequest
if err := decodeOptionalJSONBody(r, &req); err != nil {
writeError(w, http.StatusBadRequest, "bad_request", "Invalid request body")
return
}
if requiresLivePlaybackControl(name) && (session == nil || !session.HasRealtimeConnection) {
writeError(w, http.StatusConflict, "realtime_unavailable", "Realtime connection unavailable for playback session")
return
+41 -1
View File
@@ -7,6 +7,7 @@ import (
"log/slog"
"net/http"
"strconv"
"time"
"github.com/go-chi/chi/v5"
@@ -15,6 +16,7 @@ import (
"github.com/Silo-Server/silo-server/internal/downloads"
"github.com/Silo-Server/silo-server/internal/httpstream"
"github.com/Silo-Server/silo-server/internal/playback"
"github.com/Silo-Server/silo-server/internal/streamrevoke"
)
// DownloadService is the interface that the download handler depends on. A
@@ -44,7 +46,8 @@ type DownloadService interface {
// DownloadHandler handles download endpoints.
type DownloadHandler struct {
svc DownloadService
svc DownloadService
revocation *streamrevoke.Store
}
// NewDownloadHandler creates a new DownloadHandler.
@@ -52,6 +55,28 @@ func NewDownloadHandler(svc DownloadService) *DownloadHandler {
return &DownloadHandler{svc: svc}
}
// SetRevocationStore wires the stream kill switch so a per-user revocation cuts
// an IN-FLIGHT download pour. Downloads stay exempt from the live-stream cap
// (they have their own quota), but a user whose sessions were just revoked must
// not keep pulling a multi-GB file on a pre-revocation connection. Optional;
// nil keeps prior behavior.
func (h *DownloadHandler) SetRevocationStore(store *streamrevoke.Store) {
if h != nil {
h.revocation = store
}
}
// cutDownloadOnUserRevocation arms the shared in-flight cut for a download pour.
// Sessionless (downloads have no stream session), so only user-kind kills apply;
// the request entry time predates any future revocation, which is exactly the
// user-kill cutoff contract. Returns a stop func the caller must defer.
func (h *DownloadHandler) cutDownloadOnUserRevocation(w http.ResponseWriter, userID int) func() {
if h.revocation == nil {
return func() {}
}
return h.revocation.WatchAndCut(w, "", userID, time.Now())
}
// downloadRequest represents the JSON body for POST /downloads.
type downloadRequest struct {
ContentID string `json:"content_id"`
@@ -381,6 +406,13 @@ func (h *DownloadHandler) HandlePatchDownload(w http.ResponseWriter, r *http.Req
}
// HandleDownloadFile handles GET /downloads/{id}/file.
//
// This is an offline download, NOT a live stream: it creates no SessionManager
// session and no monitor record, so it is intentionally EXEMPT from the
// concurrent-stream cap and invisible to streammonitor. Downloads are bounded by
// the separate download concurrency/period quota (see the download service's
// ErrConcurrentLimitReached / ErrPeriodLimitReached). See the coverage matrix
// "downloads" note before making downloads count against the live-stream cap.
func (h *DownloadHandler) HandleDownloadFile(w http.ResponseWriter, r *http.Request) {
userID := apimw.GetUserID(r.Context())
if userID == 0 {
@@ -397,6 +429,10 @@ func (h *DownloadHandler) HandleDownloadFile(w http.ResponseWriter, r *http.Requ
return
}
// In-flight kill switch: a user revocation cuts this pour mid-transfer.
stop := h.cutDownloadOnUserRevocation(w, userID)
defer stop()
profileID, deviceID, _, _ := managedIdentity(r)
filter := requestAccessFilter(r)
// Full media downloads outlive the server's absolute WriteTimeout; roll
@@ -445,6 +481,10 @@ func (h *DownloadHandler) HandleDirectDownload(w http.ResponseWriter, r *http.Re
return
}
// In-flight kill switch: a user revocation cuts this pour mid-transfer.
stop := h.cutDownloadOnUserRevocation(w, userID)
defer stop()
filter := requestAccessFilter(r)
if err := h.svc.ServeDirect(r.Context(), w, r, userID, fileID, r.URL.Query().Get("format"), filter); err != nil {
h.writeDownloadError(w, err)
+3 -2
View File
@@ -83,8 +83,9 @@ func (w *statusWriter) Flush() {
}
}
// Unwrap lets http.ResponseController reach the underlying connection (e.g.
// for the per-response write deadlines used by streaming handlers).
// Unwrap exposes the wrapped ResponseWriter so http.ResponseController (used by
// the stream kill switch's in-flight cut via SetWriteDeadline) can reach the
// underlying socket instead of stopping at this wrapper and no-oping.
func (w *statusWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
+3 -2
View File
@@ -114,8 +114,9 @@ func (w *requestStatusWriter) Flush() {
}
}
// Unwrap lets http.ResponseController reach the underlying connection (e.g.
// for the per-response write deadlines used by streaming handlers).
// Unwrap exposes the wrapped ResponseWriter so http.ResponseController (used by
// the stream kill switch's in-flight cut via SetWriteDeadline) can reach the
// underlying socket instead of stopping at this wrapper and no-oping.
func (w *requestStatusWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
+85 -6
View File
@@ -64,6 +64,8 @@ import (
"github.com/Silo-Server/silo-server/internal/scanqueue"
"github.com/Silo-Server/silo-server/internal/secret"
"github.com/Silo-Server/silo-server/internal/sections"
"github.com/Silo-Server/silo-server/internal/streamrevoke"
"github.com/Silo-Server/silo-server/internal/streamtoken"
"github.com/Silo-Server/silo-server/internal/subtitles"
subtitleai "github.com/Silo-Server/silo-server/internal/subtitles/ai"
"github.com/Silo-Server/silo-server/internal/subtitles/opensubtitles"
@@ -167,6 +169,9 @@ type Dependencies struct {
ChapterThumbnailQueuer catalog.ChapterThumbnailQueuer
PlaybackRealtimeHub *playback.RealtimeHub
OnUserSessionsRevoked func(ctx context.Context, userID int)
// RevocationStore is the stream kill switch (may be nil). Admin terminate
// writes to it and integrated-mode serve paths consult it.
RevocationStore *streamrevoke.Store
OnServerSettingUpdated func(ctx context.Context, key, value string)
RequestServerRestart func(ctx context.Context) error
ServerRestartStatus *handlers.ServerRestartStatusTracker
@@ -1002,6 +1007,7 @@ func NewRouter(deps Dependencies) chi.Router {
playbackHandler.MarkerUpdateNotifier = playback.NewMarkerUpdateNotifier(deps.SessionMgr, realtimeHub)
subtitleAINotifier = playback.NewSubtitleReadyNotifier(deps.SessionMgr, realtimeHub)
adminPlaybackControlHandler = handlers.NewAdminPlaybackControlHandler(playbackHandler)
adminPlaybackControlHandler.SetRevocationStore(deps.RevocationStore)
if deps.DB != nil && deps.FileRepo != nil && viewerResolver != nil && deps.Config != nil && detailSvc != nil {
roomTokenService := watchtogether.NewRoomTokenService(deps.Config.Auth.JWTSecret, 24*time.Hour)
@@ -1628,6 +1634,11 @@ func NewRouter(deps Dependencies) chi.Router {
} else {
downloadHandler = handlers.NewDownloadHandler(nil)
}
// Kill switch for in-flight download pours: a per-user stream revocation
// also hangs up a download mid-transfer (downloads stay exempt from the
// live-stream cap; this only closes the "revoked user keeps pulling a
// multi-GB file on a pre-revocation connection" hole).
downloadHandler.SetRevocationStore(deps.RevocationStore)
var policyHandler *handlers.PolicyHandler
if deps.PolicySystem != nil && deps.DB != nil {
@@ -2480,8 +2491,8 @@ func NewRouter(deps Dependencies) chi.Router {
// HLS transcode delivery — no profile auth needed;
// session ID (UUID) serves as the access token, same
// pattern as /stream/{session_id}.
r.Get("/transcode/{session_id}/master.m3u8", playbackHandler.HandleGetTranscodeManifest)
r.Get("/transcode/{session_id}/segment/{name}", playbackHandler.HandleGetTranscodeSegment)
r.Get("/transcode/{session_id}/master.m3u8", guardRevocation(deps.RevocationStore, configJWTSecret(deps), playbackHandler.HandleGetTranscodeManifest))
r.Get("/transcode/{session_id}/segment/{name}", guardRevocation(deps.RevocationStore, configJWTSecret(deps), playbackHandler.HandleGetTranscodeSegment))
// Playback realtime control socket — needs auth but not profile.
r.Get("/sessions/{session_id}/control/ws", playbackHandler.HandleSessionWebSocket)
@@ -2523,10 +2534,10 @@ func NewRouter(deps Dependencies) chi.Router {
// Stream routes.
if streamHandler != nil {
r.Get("/stream/{session_id}", streamHandler.HandleStream)
r.Head("/stream/{session_id}", streamHandler.HandleStream)
r.Get("/stream/{session_id}/subtitles/{track}", streamHandler.HandleSubtitle)
r.Get("/stream/{session_id}/subtitles/{track}/fonts", streamHandler.HandleSubtitleFonts)
r.Get("/stream/{session_id}", guardRevocationCut(deps.RevocationStore, configJWTSecret(deps), streamHandler.HandleStream))
r.Head("/stream/{session_id}", guardRevocationCut(deps.RevocationStore, configJWTSecret(deps), streamHandler.HandleStream))
r.Get("/stream/{session_id}/subtitles/{track}", guardRevocation(deps.RevocationStore, configJWTSecret(deps), streamHandler.HandleSubtitle))
r.Get("/stream/{session_id}/subtitles/{track}/fonts", guardRevocation(deps.RevocationStore, configJWTSecret(deps), streamHandler.HandleSubtitleFonts))
}
// Download routes.
@@ -2928,6 +2939,7 @@ func NewRouter(deps Dependencies) chi.Router {
jwtSecret = deps.Config.Auth.JWTSecret
}
nodeHandler := handlers.NewNodeHandler(deps.NodeRepo, deps.ProxyPool, deps.TranscodePool, deps.NodeRepo, deps.EventBus, deps.RedisClient, jwtSecret)
nodeHandler.SetLocalSessionSource(deps.SessionMgr, deps.NodeID)
r.Route("/nodes", func(r chi.Router) {
r.Get("/", nodeHandler.HandleListNodes)
r.Post("/", nodeHandler.HandleCreateNode)
@@ -3484,3 +3496,70 @@ func metadataAIConfigFromServer(cfg *config.Config) metadatatranslation.Config {
OnView: cfg.MetadataAI.OnView,
}
}
// configJWTSecret returns the stream-token signing secret, or "" if unset.
func configJWTSecret(deps Dependencies) string {
if deps.Config != nil {
return deps.Config.Auth.JWTSecret
}
return ""
}
// guardRevocation wraps a native (integrated-mode) stream serve handler so the
// server refuses a revoked session/user before serving. Session id comes from
// the {session_id} URL param; user id (best-effort) from the ?st= stream token.
// This is the integrated-mode analogue of the edge nodes' IsRevoked check —
// without it a kill (admin terminate, over-cap, account revocation) would have
// no teeth on a single-node deployment. It refuses new/next requests (stopping
// HLS on its next segment and refusing reconnects); cutting an in-flight native
// direct-play pour is a follow-up.
// streamRequestIdentity extracts the best-effort owner user id and the
// credential-issue time for a native serve request. With a valid ?st= stream
// token those are the token's uid and iat — the keys the user-kill cutoff in
// streamrevoke compares (a token minted before a user revocation is killed; one
// minted after re-auth plays). Without a token (legacy/bare URL or signing
// disabled) these routes still run after API auth, so the authenticated user id
// plus the request entry time apply: an entry request postdates any existing
// user kill (fresh auth ⇒ allowed) while an in-flight pour predates a future
// one (cut by WatchAndCut).
func streamRequestIdentity(r *http.Request, secret string) (int, time.Time) {
if tok := r.URL.Query().Get("st"); tok != "" && secret != "" {
if claims, err := streamtoken.Verify(tok, secret); err == nil {
return claims.UserID, claims.IssuedTime()
}
}
return apimw.GetUserID(r.Context()), time.Now()
}
func guardRevocation(store *streamrevoke.Store, secret string, next http.HandlerFunc) http.HandlerFunc {
if store == nil {
return next
}
return func(w http.ResponseWriter, r *http.Request) {
userID, startedAt := streamRequestIdentity(r, secret)
if store.Refuse(w, chi.URLParam(r, "session_id"), userID, startedAt) {
return
}
next(w, r)
}
}
// guardRevocationCut is guardRevocation plus an in-flight connection cut, for the
// long single-GET pours (native direct-play/remux on /stream). Transcode segment
// routes use plain guardRevocation — per-segment refusal already stops HLS within
// one segment, so a per-segment watcher goroutine would be wasteful.
func guardRevocationCut(store *streamrevoke.Store, secret string, next http.HandlerFunc) http.HandlerFunc {
if store == nil {
return next
}
return func(w http.ResponseWriter, r *http.Request) {
sessionID := chi.URLParam(r, "session_id")
userID, startedAt := streamRequestIdentity(r, secret)
if store.Refuse(w, sessionID, userID, startedAt) {
return
}
stop := store.WatchAndCut(w, sessionID, userID, startedAt)
defer stop()
next(w, r)
}
}
+32
View File
@@ -25,6 +25,7 @@ import (
"github.com/Silo-Server/silo-server/internal/models"
"github.com/Silo-Server/silo-server/internal/nodepool"
"github.com/Silo-Server/silo-server/internal/playback"
"github.com/Silo-Server/silo-server/internal/streamrevoke"
"github.com/Silo-Server/silo-server/internal/streamtoken"
"github.com/Silo-Server/silo-server/internal/subtitles"
"github.com/Silo-Server/silo-server/internal/transcodenode"
@@ -206,6 +207,11 @@ type PlaybackHandler struct {
// server-authoritative store instead (see internal/noderecipe). Optional
// (nil disables it — integrated/no-node deployments need no handoff).
RecipeNodeStore recipeNodePutter
// Revocation is the shared stream kill switch. jellycompat local serving (the
// integrated path, and multi-node local-transcode fallback) must consult it so
// an admin terminate / over-cap / account revocation actually stops the bytes —
// the multi-node redirect path is already guarded by the proxy. Optional.
Revocation *streamrevoke.Store
}
// recipeNodePutter persists and removes a remote transcode's reconstruction
@@ -372,6 +378,20 @@ func (h *PlaybackHandler) buildProxyRedirectURL(
AudioTrackIndex: audioTrackIndex,
TranscodeNode: transcodeNodeURL,
DVProfile: file.PrimaryDVProfile(),
Origin: playback.OriginJellyfin,
}
// Carry the owner (uid/pid/mfid) so the edge tracker attributes this stream to
// the real user — without it jellycompat streams land under user 0 in the
// monitoring picture, breaking per-user over-cap enforcement and user-key kills
// at the edge. Native tokens already carry the owner; keep parity. Also carry
// the client name so the edge record identifies the viewer/app.
if h.sessionMgr != nil {
if s, err := h.sessionMgr.GetSession(upstreamSessionID); err == nil && s != nil {
claims.UserID = s.UserID
claims.ProfileID = s.ProfileID
claims.MediaFileID = s.MediaFileID
claims.ClientName = s.ClientName
}
}
token, err := streamtoken.Sign(claims, h.JWTSecret, 24*time.Hour)
if err != nil {
@@ -454,6 +474,18 @@ func (h *PlaybackHandler) startRemoteTranscode(
if source.TranscodeAudio {
reqBody.TargetCodecVideo = "copy"
}
// Attribute the node's live-session record to the real owner + route + client
// so it is not an ownerless (user 0), route-less record when no proxy record
// fronts it. Mirrors the redirect token's owner-carry.
reqBody.Route = playback.OriginJellyfin
if h.sessionMgr != nil {
if s, err := h.sessionMgr.GetSession(upstreamSessionID); err == nil && s != nil {
reqBody.AuthUserID = s.UserID
reqBody.ProfileID = s.ProfileID
reqBody.MediaFileID = s.MediaFileID
reqBody.ClientName = s.ClientName
}
}
body, err := json.Marshal(reqBody)
if err != nil {
+8
View File
@@ -96,6 +96,14 @@ func (w *compatImageProxyTagResponseWriter) finish() {
}
}
// Unwrap exposes the wrapped ResponseWriter so http.ResponseController (used by
// the stream kill switch's in-flight cut via SetWriteDeadline) can reach the
// underlying socket instead of stopping at this wrapper and no-oping. Video
// stream responses are non-JSON, so they pass through this writer untouched.
func (w *compatImageProxyTagResponseWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
func isJSONResponse(contentType string) bool {
contentType = strings.ToLower(strings.TrimSpace(contentType))
return contentType == "application/json" || strings.HasPrefix(contentType, "application/json;")
+6 -2
View File
@@ -50,7 +50,9 @@ func (w *loggingResponseWriter) Flush() {
}
}
// Unwrap returns the underlying ResponseWriter for http.ResponseController.
// Unwrap exposes the wrapped ResponseWriter so http.ResponseController (used by
// the stream kill switch's in-flight cut via SetWriteDeadline) can reach the
// underlying socket instead of stopping at this wrapper and no-oping.
func (w *loggingResponseWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
@@ -156,7 +158,9 @@ func (w *debugResponseWriter) Flush() {
}
}
// Unwrap returns the underlying ResponseWriter for http.ResponseController.
// Unwrap exposes the wrapped ResponseWriter so http.ResponseController (used by
// the stream kill switch's in-flight cut via SetWriteDeadline) can reach the
// underlying socket instead of stopping at this wrapper and no-oping.
func (w *debugResponseWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
+1
View File
@@ -108,6 +108,7 @@ func NewRouter(deps Dependencies) chi.Router {
}
playbackHandler.NodePlanner = deps.NodePlanner
playbackHandler.JWTSecret = deps.JWTSecret
playbackHandler.Revocation = deps.RevocationStore
// Compat transcode reconstruct is driven by the recipe carried in the durable
// compat playback store (jellycompat_playback_sessions); no separate native
// recipe table is needed.
+4
View File
@@ -17,6 +17,7 @@ import (
"github.com/Silo-Server/silo-server/internal/recommendations"
"github.com/Silo-Server/silo-server/internal/scantrigger"
"github.com/Silo-Server/silo-server/internal/secret"
"github.com/Silo-Server/silo-server/internal/streamrevoke"
"github.com/Silo-Server/silo-server/internal/subtitles"
"github.com/Silo-Server/silo-server/internal/userstore"
"github.com/Silo-Server/silo-server/internal/watchstate"
@@ -30,6 +31,9 @@ type Dependencies struct {
// periodic orphan-transcode sweep so it stops on shutdown; nil (tests) makes
// the sweep a single boot-time run instead of a long-lived ticker.
AppContext context.Context
// RevocationStore is the shared stream kill switch consulted by local serving.
RevocationStore *streamrevoke.Store
// LiveConfig returns the current hot-reloaded config. May be nil (tests,
// worker modes); read through CurrentConfig(), which falls back to Config.
LiveConfig func() *config.Config
+95 -1
View File
@@ -87,12 +87,35 @@ func (h *PlaybackHandler) HandleVideoStream(w http.ResponseWriter, r *http.Reque
}
}
// Kill switch, BEFORE ensureUpstreamPlayback: a killed upstream session must
// be refused while it is still identifiable. ensureUpstreamPlayback revives a
// lost session — or, when the revoked one is no longer reconstructable, mints
// a REPLACEMENT with a fresh id that the post-ensure check below would wave
// through, letting a client dodge a session kill by simply re-hitting the same
// stream URL. (User-kind kills don't need this: they match by user id either
// side of the ensure.)
if h.Revocation != nil && playSession.UpstreamSessionID != "" &&
h.Revocation.Refuse(w, playSession.UpstreamSessionID, session.StreamAppUserID, time.Now()) {
return
}
playSession, err = h.ensureUpstreamPlayback(r.Context(), session, playSession.ID, *source, method)
if err != nil {
writeCompatUpstreamError(w, err)
return
}
// Kill switch (local serving): refuse a revoked session/user before serving.
// The multi-node redirect path below is guarded by the proxy; this covers the
// integrated / local-fallback path that the proxy never sees. The request
// entry time is the credential time for the user-kill cutoff: a compat
// request only gets here on a live compat login, and user revocation deletes
// all compat logins, so reaching this line post-revocation means fresh auth.
if h.Revocation != nil && playSession.UpstreamSessionID != "" &&
h.Revocation.Refuse(w, playSession.UpstreamSessionID, session.StreamAppUserID, time.Now()) {
return
}
if h.fileResolver == nil {
writeError(w, http.StatusInternalServerError, "ServerError", "File resolver not available")
return
@@ -127,6 +150,14 @@ func (h *PlaybackHandler) HandleVideoStream(w http.ResponseWriter, r *http.Reque
}
}
// In-flight cut: hang up a long local direct-play/remux pour the moment it is
// revoked. The Refuse check above only covers new/reconnect requests; a single
// long GET needs the connection cut to stop mid-stream.
if h.Revocation != nil && playSession.UpstreamSessionID != "" {
stop := h.Revocation.WatchAndCut(w, playSession.UpstreamSessionID, session.StreamAppUserID, time.Now())
defer stop()
}
switch method {
case "remux":
audioTrackIndex := -1
@@ -144,6 +175,19 @@ func (h *PlaybackHandler) HandleVideoStream(w http.ResponseWriter, r *http.Reque
// load-bearing for Infuse: it refuses Direct Play (Static=true streaming)
// for items it believes it cannot download, so the flag must stay true and
// this route must exist.
// HandleDownload serves a full file for offline/download clients (e.g. Infuse
// with CanDownload). By design it is NOT a live stream: it creates no
// SessionManager session and no monitor record, so it is EXEMPT from the
// concurrent-stream cap and invisible to streammonitor. NOTE: unlike the native
// /downloads/{id}/file route, this route is NOT covered by the download
// concurrency/period quota either — it requires no download row, only compat
// auth. Access control is compat authentication plus the library access filter;
// a per-user stream revocation cuts an in-flight pour (below) and, because the
// same hook deletes every compat login, refuses reconnects until re-auth.
// Bringing this route under a real download quota (Infuse-compatible: plain
// GETs, Range requests, no download rows) is an open follow-up in the coverage
// matrix; if downloads should ever count against the live-stream cap instead,
// register a tracked session here — but that conflates two distinct limits.
func (h *PlaybackHandler) HandleDownload(w http.ResponseWriter, r *http.Request) {
session := SessionFromContext(r.Context())
if session == nil {
@@ -184,6 +228,14 @@ func (h *PlaybackHandler) HandleDownload(w http.ResponseWriter, r *http.Request)
return
}
// In-flight kill switch: sessionless pour, so only user-kind kills apply;
// the entry time predates any future revocation (the user-kill cutoff), so a
// revocation issued mid-transfer hangs this connection up.
if h.Revocation != nil {
stop := h.Revocation.WatchAndCut(w, "", session.StreamAppUserID, time.Now())
defer stop()
}
w.Header().Set("Content-Disposition", "attachment; filename*=UTF-8''"+url.PathEscape(filepath.Base(file.FilePath)))
_ = playback.ServeDirectPlay(w, r, file.FilePath)
}
@@ -210,6 +262,16 @@ func (h *PlaybackHandler) HandleMasterManifest(w http.ResponseWriter, r *http.Re
return
}
// Kill switch: refuse a revoked session before the ensure/start machinery
// below can revive its transcode. Without this a killed session's player
// polling the manifest keeps ffmpeg alive (and re-spawns it after a restart)
// even though every segment request is refused — the same resurrection hole
// the transcode node's reconstruct guard closes on the edge.
if h.Revocation != nil && playSession.UpstreamSessionID != "" &&
h.Revocation.Refuse(w, playSession.UpstreamSessionID, session.StreamAppUserID, time.Now()) {
return
}
source := findMediaSource(playSession, firstNonEmpty(r.URL.Query().Get("MediaSourceId"), r.URL.Query().Get("mediaSourceId")))
if source == nil {
writeError(w, http.StatusBadRequest, "BadRequest", "Media source is required")
@@ -319,6 +381,12 @@ func (h *PlaybackHandler) HandleHLSManifest(w http.ResponseWriter, r *http.Reque
writeError(w, http.StatusNotFound, "NotFound", "Playback session not found")
return
}
// Kill switch: refuse before ensureTranscodeManifest can keep/restart the
// killed session's ffmpeg (see HandleMasterManifest).
if h.Revocation != nil && playSession.UpstreamSessionID != "" &&
h.Revocation.Refuse(w, playSession.UpstreamSessionID, session.StreamAppUserID, time.Now()) {
return
}
source := firstMediaSource(playSession)
if mediaSourceID := firstNonEmpty(r.URL.Query().Get("MediaSourceId"), r.URL.Query().Get("mediaSourceId")); mediaSourceID != "" {
source = findMediaSource(playSession, mediaSourceID)
@@ -372,6 +440,12 @@ func (h *PlaybackHandler) HandleHLSSegment(w http.ResponseWriter, r *http.Reques
return
}
// Kill switch (local serving): refuse a revoked session/user on every segment.
// Request entry time = credential time (compat auth is live, see HandleVideoStream).
if h.Revocation != nil && h.Revocation.Refuse(w, playSession.UpstreamSessionID, session.StreamAppUserID, time.Now()) {
return
}
name := chiURLParam(r, "segmentId")
ext := chiURLParam(r, "segmentContainer")
@@ -531,6 +605,18 @@ func (h *PlaybackHandler) HandleHLSSegment(w http.ResponseWriter, r *http.Reques
transcodeSession.ReportSegmentDownloaded(segNum)
}
// Hold a transport marker for the whole segment pour so a compat HLS transcode
// stays visible via server-observed liveness rather than the client's progress
// POST: a client that keeps pulling segments but withholds progress reports
// must still be counted (the hidden-stream defense). Mirrors the direct/remux
// path above and the native transcode segment handler.
if h.sessionMgr != nil && playSession.UpstreamSessionID != "" {
if err := h.sessionMgr.BeginTransport(playSession.UpstreamSessionID); err == nil {
upstreamSessionID := playSession.UpstreamSessionID
defer func() { _ = h.sessionMgr.EndTransport(upstreamSessionID) }()
}
}
http.ServeFile(w, r, segmentPath)
}
@@ -559,12 +645,20 @@ func (h *PlaybackHandler) HandleSubtitleStream(w http.ResponseWriter, r *http.Re
return
}
_, source, err := h.resolvePlaybackRoute(r, session, chiURLParam(r, "routeMediaSourceId"), chiURLParam(r, "routeMediaSourceId"))
playSession, source, err := h.resolvePlaybackRoute(r, session, chiURLParam(r, "routeMediaSourceId"), chiURLParam(r, "routeMediaSourceId"))
if err != nil || source == nil {
writeError(w, http.StatusNotFound, "NotFound", "Playback session not found")
return
}
// Kill switch: a revoked session must not keep pulling subtitle bytes (or
// triggering server-side ffmpeg subtitle extraction below). Mirrors the
// guarded native subtitle routes.
if h.Revocation != nil && playSession != nil && playSession.UpstreamSessionID != "" &&
h.Revocation.Refuse(w, playSession.UpstreamSessionID, session.StreamAppUserID, time.Now()) {
return
}
if h.fileResolver == nil {
writeError(w, http.StatusInternalServerError, "ServerError", "File resolver not available")
return
+26 -3
View File
@@ -1,6 +1,7 @@
package proxy
import (
"io"
"net/http"
"sync"
"time"
@@ -57,9 +58,11 @@ func (m *egressMeter) RateKbps() int {
return int(total * 8 / 1000 / meterWindowSeconds)
}
// meteredResponseWriter counts every byte written to the client.
// Embedding the interface intentionally hides optimizations like
// io.ReaderFrom so all writes flow through Write.
// meteredResponseWriter counts every byte written to the client. Embedding the
// interface hides optimizations like io.ReaderFrom, so ReadFrom is forwarded
// explicitly below — every stream handler runs inside meterEgress, so without
// that forwarding NO proxied pour could ever reach the kernel sendfile path,
// regardless of what the inner writers (sessionByteWriter) forward.
type meteredResponseWriter struct {
http.ResponseWriter
meter *egressMeter
@@ -71,12 +74,32 @@ func (w *meteredResponseWriter) Write(b []byte) (int, error) {
return n, err
}
// ReadFrom preserves the underlying ResponseWriter's sendfile fast path
// (disk→socket inside the kernel) while still metering the bytes: the count
// only needs the total, which sendfile reports on return. Without a ReaderFrom
// underneath it falls back to a plain copy through Write (which meters), via
// writeOnly so io.Copy cannot re-enter this method and recurse.
func (w *meteredResponseWriter) ReadFrom(src io.Reader) (int64, error) {
if rf, ok := w.ResponseWriter.(io.ReaderFrom); ok {
n, err := rf.ReadFrom(src)
w.meter.Add(n)
return n, err
}
return io.Copy(writeOnly{w}, src)
}
func (w *meteredResponseWriter) Flush() {
if f, ok := w.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
}
// Unwrap lets http.NewResponseController reach the underlying ResponseWriter so
// SetWriteDeadline (used by the revocation connection-cut) can find the socket.
func (w *meteredResponseWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
// meterEgress wraps stream handlers so their responses count toward the
// node's measured egress bandwidth.
func (s *Server) meterEgress(next http.Handler) http.Handler {
+161 -9
View File
@@ -5,6 +5,7 @@ import (
"fmt"
"io"
"log/slog"
"net"
"net/http"
"strconv"
"strings"
@@ -19,6 +20,15 @@ import (
"github.com/Silo-Server/silo-server/internal/streamtoken"
)
// revocationStore is the subset of *streamrevoke.Store the edge consults to
// enforce kills. IsRevoked is a pure in-memory lookup (no I/O) so it is safe on
// the per-request hot path. Nil disables enforcement.
type revocationStore interface {
IsRevoked(sessionID string, userID int, startedAt time.Time) bool
Refuse(w http.ResponseWriter, sessionID string, userID int, startedAt time.Time) bool
WatchAndCut(w http.ResponseWriter, sessionID string, userID int, startedAt time.Time) func()
}
// Server is the HTTP handler for proxy mode.
type Server struct {
watcher *nodeconfig.Watcher
@@ -28,6 +38,15 @@ type Server struct {
// subCache stores full-track PGS (.sup) extracts under the transcode dir
// so repeat selections skip the whole-file ffmpeg demux.
subCache *playback.SubtitleCache
revocation revocationStore
}
// SetRevocationStore wires the kill-switch the edge consults per request. The
// store's cache is kept current out-of-band (pub/sub + poll), so the hot-path
// check is a local map read.
func (s *Server) SetRevocationStore(store revocationStore) {
s.revocation = store
}
// NewServer creates a new proxy server backed by a config watcher and session
@@ -137,20 +156,119 @@ func (s *Server) verifyToken(w http.ResponseWriter, r *http.Request) *streamtoke
http.Error(w, "unauthorized", http.StatusUnauthorized)
return nil
}
// Kill switch: a revoked session/user is refused here, on every request. For
// chunked (HLS) playback this stops the stream on its next segment fetch and
// refuses reconnects; long direct-play/remux pours are additionally cut mid-
// stream by cutOnRevocation.
// The token's iat is the credential-issue time the user-kill cutoff compares
// against: edge requests carry no fresh auth, only this token, so a user
// revocation kills exactly the tokens minted before it.
if s.revocation != nil && s.revocation.Refuse(w, claims.SessionID, claims.UserID, claims.IssuedTime()) {
return nil
}
return claims
}
// cutOnRevocation watches a long-lived pour (direct play / remux) and hangs up
// the socket the moment the session/user is revoked, via the shared store helper.
// Returns a stop func to cancel the watcher when the request finishes normally.
func (s *Server) cutOnRevocation(w http.ResponseWriter, claims *streamtoken.Claims) func() {
if s.revocation == nil {
return func() {}
}
return s.revocation.WatchAndCut(w, claims.SessionID, claims.UserID, claims.IssuedTime())
}
func (s *Server) handleDirectPlay(w http.ResponseWriter, r *http.Request) {
claims := s.verifyToken(w, r)
if claims == nil {
return
}
info := sessionInfo(s.tracker, claims, "direct_play")
info := sessionInfo(s.tracker, claims, "direct_play", edgeClientIP(r))
s.tracker.Track(r.Context(), info)
defer s.tracker.Remove(r.Context(), claims.SessionID)
http.ServeFile(w, r, claims.MediaPath)
sw := &sessionByteWriter{ResponseWriter: w, tracker: s.tracker, sessionID: claims.SessionID}
defer sw.flush()
stop := s.cutOnRevocation(sw, claims)
defer stop()
http.ServeFile(sw, r, claims.MediaPath)
}
// sessionByteWriter attributes served bytes to a session so LastServedAt advances
// during long direct-play/remux pours (authoritative liveness). Bytes are
// flushed to the tracker in coarse chunks to avoid per-write lock churn. Unwrap
// lets the revocation connection-cut reach the socket through this layer.
type sessionByteWriter struct {
http.ResponseWriter
tracker *nodesessions.Tracker
sessionID string
acc int64
}
func (w *sessionByteWriter) Write(b []byte) (int, error) {
n, err := w.ResponseWriter.Write(b)
w.account(int64(n))
return n, err
}
// account tallies served bytes, flushing to the tracker in coarse ~1 MiB chunks
// to avoid per-write lock churn. Shared by Write and the ReadFrom fast path.
func (w *sessionByteWriter) account(n int64) {
if n <= 0 {
return
}
w.acc += n
if w.acc >= 1<<20 { // flush every ~1 MiB
w.tracker.AddBytes(w.sessionID, w.acc)
w.acc = 0
}
}
// ReadFrom preserves the underlying writer's sendfile fast path while still
// attributing served bytes. http.ServeFile pours an *os.File through the
// ResponseWriter's io.ReaderFrom (sendfile: disk->socket inside the kernel, no
// userspace copy) — but only if it can see that method. Wrapping the writer for
// byte counting would hide it and force every direct-play/remux byte through
// userspace; forwarding here keeps zero-copy AND the count. Liveness does not
// depend on this coarse count: the session stays live from Track to Remove
// (tracker re-SETs it, never idle-prunes an open pour), so a single sendfile
// call that only tallies on return still stays visible for the whole pour.
func (w *sessionByteWriter) ReadFrom(src io.Reader) (int64, error) {
if rf, ok := w.ResponseWriter.(io.ReaderFrom); ok {
n, err := rf.ReadFrom(src)
w.account(n)
return n, err
}
// No sendfile fast path underneath. Copy manually — but NOT via io.Copy(w,
// src), which would re-detect this very ReadFrom and recurse forever.
// writeOnly exposes only Write, so io.Copy falls back to the Write loop
// (which accounts bytes).
return io.Copy(writeOnly{w}, src)
}
// writeOnly wraps an io.Writer to expose ONLY Write, hiding any ReadFrom so
// io.Copy cannot re-enter sessionByteWriter.ReadFrom.
type writeOnly struct{ io.Writer }
func (w *sessionByteWriter) flush() {
if w.acc > 0 {
w.tracker.AddBytes(w.sessionID, w.acc)
w.acc = 0
}
}
func (w *sessionByteWriter) Flush() {
if f, ok := w.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
}
func (w *sessionByteWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
func (s *Server) handleRemux(w http.ResponseWriter, r *http.Request) {
@@ -159,10 +277,16 @@ func (s *Server) handleRemux(w http.ResponseWriter, r *http.Request) {
return
}
info := sessionInfo(s.tracker, claims, "remux")
info := sessionInfo(s.tracker, claims, "remux", edgeClientIP(r))
s.tracker.Track(r.Context(), info)
defer s.tracker.Remove(r.Context(), claims.SessionID)
sw := &sessionByteWriter{ResponseWriter: w, tracker: s.tracker, sessionID: claims.SessionID}
defer sw.flush()
stop := s.cutOnRevocation(sw, claims)
defer stop()
seekSeconds := 0.0
if seekStr := r.URL.Query().Get("seek"); seekStr != "" {
if v, err := strconv.ParseFloat(seekStr, 64); err == nil {
@@ -171,8 +295,9 @@ func (s *Server) handleRemux(w http.ResponseWriter, r *http.Request) {
}
// Honor the Dolby Vision mode frozen in the token (empty decodes as the
// legacy auto behavior for old tokens), mirroring how the integrated
// server's stream handler serves the same claims.
_ = playback.ServeRemuxWithDVMode(w, r, claims.MediaPath, "mp4", seekSeconds, claims.TranscodeAudio, claims.AudioTrackIndex, claims.DVProfile, playback.RemuxDVMode(claims.RemuxDVMode), s.watcher.Config().Playback.FFmpegPath)
// server's stream handler serves the same claims. Serve through sw so the
// bytes are metered and the revocation cut can reach this response.
_ = playback.ServeRemuxWithDVMode(sw, r, claims.MediaPath, "mp4", seekSeconds, claims.TranscodeAudio, claims.AudioTrackIndex, claims.DVProfile, playback.RemuxDVMode(claims.RemuxDVMode), s.watcher.Config().Playback.FFmpegPath)
}
func (s *Server) handleTranscodeManifest(w http.ResponseWriter, r *http.Request) {
@@ -209,17 +334,24 @@ func transcodeTransportIDFromClaims(claims *streamtoken.Claims) string {
// short manifest/segment requests, so the session is tracked by recent
// activity instead of request lifetime.
func (s *Server) touchTranscodeSession(r *http.Request, claims *streamtoken.Claims) {
s.tracker.Touch(r.Context(), sessionInfo(s.tracker, claims, "transcode"))
s.tracker.Touch(r.Context(), sessionInfo(s.tracker, claims, "transcode", edgeClientIP(r)))
}
// sessionInfo builds the node-session tracker record for a verified token,
// copying the numeric ownership keys the node-session tracker needs.
func sessionInfo(tr *nodesessions.Tracker, claims *streamtoken.Claims, kind string) nodesessions.SessionInfo {
// copying the numeric ownership keys plus the monitoring attribution (route +
// client identity) the first-class monitor view needs. clientIP is the connecting
// address observed at the edge (best-effort; the client reaches the proxy
// directly). Route/ClientName come from the token so the edge — which never sees
// the originating API path — can still tag native vs jellycompat.
func sessionInfo(tr *nodesessions.Tracker, claims *streamtoken.Claims, kind, clientIP string) nodesessions.SessionInfo {
return nodesessions.SessionInfo{
SessionID: claims.SessionID,
NodeURL: tr.NodeURL(),
NodeName: tr.NodeName(),
Type: kind,
Route: claims.Origin,
ClientIP: clientIP,
ClientName: claims.ClientName,
StartedAt: time.Now().UTC().Format(time.RFC3339),
AuthUserID: claims.UserID,
ProfileID: claims.ProfileID,
@@ -227,6 +359,15 @@ func sessionInfo(tr *nodesessions.Tracker, claims *streamtoken.Claims, kind stri
}
}
// edgeClientIP extracts the connecting client's IP from the request, best-effort.
func edgeClientIP(r *http.Request) string {
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return r.RemoteAddr
}
return host
}
func (s *Server) handleSubtitle(w http.ResponseWriter, r *http.Request) {
claims := s.verifyToken(w, r)
if claims == nil {
@@ -370,7 +511,18 @@ func (s *Server) proxyToTranscodeNode(w http.ResponseWriter, r *http.Request, cl
}
}
w.WriteHeader(resp.StatusCode)
io.Copy(w, resp.Body)
// Attribute served bytes to the session for authoritative monitoring —
// incrementally (~1 MiB granularity), not in one post-copy tally: a slow
// segment drain that outlives the 60s record TTL would otherwise go
// invisible mid-pour and then post bytes to a record that no longer exists.
// Manifest bytes are tiny; segment bytes are the real signal. Best-effort,
// never gates. writeOnly forces the per-chunk Write path: sw.ReadFrom would
// tally only once on return, which is the single-post-copy behavior this
// replaces (there is no file to sendfile here — the source is the node's
// HTTP response body).
sw := &sessionByteWriter{ResponseWriter: w, tracker: s.tracker, sessionID: claims.SessionID}
defer sw.flush()
_, _ = io.Copy(writeOnly{sw}, resp.Body)
}
func (s *Server) handleForceReload(w http.ResponseWriter, r *http.Request) {
+172
View File
@@ -0,0 +1,172 @@
package proxy
import (
"bytes"
"io"
"net/http"
"strings"
"testing"
"github.com/Silo-Server/silo-server/internal/nodesessions"
)
// rfResponseWriter is the sendfile-capable writer shape: an http.ResponseWriter
// that also implements io.ReaderFrom. It records whether ReadFrom was taken.
type rfResponseWriter struct {
buf bytes.Buffer
hdr http.Header
usedRF bool
}
func (w *rfResponseWriter) Header() http.Header {
if w.hdr == nil {
w.hdr = http.Header{}
}
return w.hdr
}
func (w *rfResponseWriter) Write(b []byte) (int, error) { return w.buf.Write(b) }
func (w *rfResponseWriter) WriteHeader(int) {}
func (w *rfResponseWriter) ReadFrom(src io.Reader) (int64, error) {
w.usedRF = true
return w.buf.ReadFrom(src)
}
// plainResponseWriter implements ONLY http.ResponseWriter (no ReadFrom), forcing
// sessionByteWriter.ReadFrom down its manual-copy fallback.
type plainResponseWriter struct {
buf bytes.Buffer
hdr http.Header
}
func (w *plainResponseWriter) Header() http.Header {
if w.hdr == nil {
w.hdr = http.Header{}
}
return w.hdr
}
func (w *plainResponseWriter) Write(b []byte) (int, error) { return w.buf.Write(b) }
func (w *plainResponseWriter) WriteHeader(int) {}
// TestSessionByteWriterReadFromFastPath verifies that when the underlying writer
// supports sendfile (io.ReaderFrom), ReadFrom forwards to it (preserving
// zero-copy) and still attributes the served bytes.
func TestSessionByteWriterReadFromFastPath(t *testing.T) {
under := &rfResponseWriter{}
sw := &sessionByteWriter{ResponseWriter: under, sessionID: "s"}
const payload = "the quick brown fox jumps"
n, err := sw.ReadFrom(strings.NewReader(payload))
if err != nil {
t.Fatalf("ReadFrom: %v", err)
}
if !under.usedRF {
t.Fatal("expected the underlying io.ReaderFrom (sendfile) fast path to be used")
}
if n != int64(len(payload)) {
t.Fatalf("bytes forwarded = %d, want %d", n, len(payload))
}
if sw.acc != int64(len(payload)) {
t.Fatalf("accounted bytes = %d, want %d", sw.acc, len(payload))
}
if got := under.buf.String(); got != payload {
t.Fatalf("served body = %q, want %q", got, payload)
}
}
// TestMeteredWriterChainPreservesSendfile proves the PRODUCTION writer chain —
// sessionByteWriter over meteredResponseWriter (every stream route runs inside
// meterEgress) — still reaches the underlying sendfile fast path AND meters the
// bytes. Regression guard: meteredResponseWriter used to hide io.ReaderFrom,
// which silently forced every proxied pour through a userspace copy no matter
// what the inner writer forwarded, making sessionByteWriter's fast path dead
// code on real requests.
func TestMeteredWriterChainPreservesSendfile(t *testing.T) {
under := &rfResponseWriter{}
meter := newEgressMeter()
mw := &meteredResponseWriter{ResponseWriter: under, meter: meter}
sw := &sessionByteWriter{ResponseWriter: mw, sessionID: "s"}
const payload = "kernel to socket, no userspace detours"
n, err := sw.ReadFrom(strings.NewReader(payload))
if err != nil {
t.Fatalf("ReadFrom: %v", err)
}
if !under.usedRF {
t.Fatal("expected sendfile fast path through the full metered chain")
}
if n != int64(len(payload)) {
t.Fatalf("bytes forwarded = %d, want %d", n, len(payload))
}
if sw.acc != int64(len(payload)) {
t.Fatalf("session-accounted bytes = %d, want %d", sw.acc, len(payload))
}
var metered int64
meter.mu.Lock()
for _, b := range meter.buckets {
metered += b
}
meter.mu.Unlock()
if metered != int64(len(payload)) {
t.Fatalf("metered bytes = %d, want %d", metered, len(payload))
}
if got := under.buf.String(); got != payload {
t.Fatalf("served body = %q, want %q", got, payload)
}
}
// TestSessionByteWriterAccountFlushBranch exercises account()'s coarse-flush
// branch (w.acc >= 1<<20) that the short-payload tests never reach: a >=1MiB
// pour must flush the accumulator to the tracker and reset acc to 0. It also
// pins that the flush is safe against a real *nodesessions.Tracker — the other
// tests leave tracker nil, which only stays safe while the branch is untaken. A
// Redis-less tracker makes AddBytes a no-op, so the branch runs without Redis.
func TestSessionByteWriterAccountFlushBranch(t *testing.T) {
tracker := nodesessions.NewTracker(nil, "http://node", "node", "proxy")
under := &rfResponseWriter{}
sw := &sessionByteWriter{ResponseWriter: under, tracker: tracker, sessionID: "s"}
const oneMiB = 1 << 20
payload := strings.Repeat("x", oneMiB)
n, err := sw.ReadFrom(strings.NewReader(payload))
if err != nil {
t.Fatalf("ReadFrom: %v", err)
}
if !under.usedRF {
t.Fatal("expected the sendfile fast path even for a large payload")
}
if n != int64(oneMiB) {
t.Fatalf("bytes forwarded = %d, want %d", n, oneMiB)
}
// The >=1MiB pour must have tripped the flush branch, resetting the accumulator.
if sw.acc != 0 {
t.Fatalf("accumulator = %d, want 0 (coarse-flush branch not taken)", sw.acc)
}
if got := under.buf.Len(); got != oneMiB {
t.Fatalf("served bytes = %d, want %d", got, oneMiB)
}
}
// TestSessionByteWriterReadFromFallbackNoRecursion verifies the fallback path
// (underlying writer has no sendfile support) copies via Write without
// re-entering ReadFrom. A naive io.Copy(sw, src) fallback would re-detect this
// ReadFrom and recurse until the stack overflows, so simply returning here — and
// counting the bytes — proves the guard works.
func TestSessionByteWriterReadFromFallbackNoRecursion(t *testing.T) {
under := &plainResponseWriter{}
sw := &sessionByteWriter{ResponseWriter: under, sessionID: "s"}
const payload = "fallback path must not recurse"
n, err := sw.ReadFrom(strings.NewReader(payload))
if err != nil {
t.Fatalf("ReadFrom: %v", err)
}
if n != int64(len(payload)) {
t.Fatalf("bytes copied = %d, want %d", n, len(payload))
}
if sw.acc != int64(len(payload)) {
t.Fatalf("accounted bytes = %d, want %d", sw.acc, len(payload))
}
if got := under.buf.String(); got != payload {
t.Fatalf("served body = %q, want %q", got, payload)
}
}
+158
View File
@@ -0,0 +1,158 @@
// Package streamenforcer is the async decision loop ("the brain") of the
// monitor-and-kill design. It runs on central only, off the hot path: it reads
// the authoritative live-streams snapshot, compares each user's live count to
// that user's limit, and issues revocations for the over-cap sessions. The
// revocation kill switch (internal/streamrevoke) then stops them at the edge
// within one propagation/poll interval.
//
// Every enforcement reason (over-cap here, admin terminate and account
// revocation elsewhere) collapses to the same action: write a revocation. This
// package owns only the over-cap rule; other reasons call the revoker directly.
package streamenforcer
import (
"context"
"log/slog"
"sort"
"time"
"github.com/Silo-Server/silo-server/internal/streammonitor"
)
// DefaultInterval is how often the enforcer evaluates the live picture. The
// ~120s enforcement budget = this interval + the revocation propagation/poll.
const DefaultInterval = 30 * time.Second
// revocationTTL is how long an over-cap revocation lasts. It is deliberately
// short: the enforcer re-evaluates every DefaultInterval, so a persistent abuser
// is re-revoked long before this lapses (staying dead), while a transient
// over-count — e.g. a ghost session lingering in the monitor next to a fresh
// reconnect — self-heals within this window instead of banning the reconnect for
// 24h. Must comfortably exceed DefaultInterval so re-revocation has slack.
const revocationTTL = 5 * time.Minute
// Revoker is the subset of *streamrevoke.Store the enforcer needs. It uses the
// TTL-scoped variant so an over-cap kill self-heals (see revocationTTL).
type Revoker interface {
RevokeSessionFor(ctx context.Context, sessionID, reason string, ttl time.Duration) error
}
// LimitFunc returns the maximum concurrent streams allowed for a user. A return
// of <= 0 means "unlimited" (no enforcement for that user), matching the
// SessionManager's convention. An error means the limit is currently unknown;
// the enforcer fails OPEN (does not kill) so a limit-lookup blip never
// terminates legitimate playback.
type LimitFunc func(ctx context.Context, userID int) (maxStreams int, err error)
// Enforcer periodically trims over-cap streams.
type Enforcer struct {
source streammonitor.Source
limits LimitFunc
revoker Revoker
interval time.Duration
now func() time.Time
}
// New builds an enforcer. interval <= 0 uses DefaultInterval.
func New(source streammonitor.Source, limits LimitFunc, revoker Revoker, interval time.Duration) *Enforcer {
if interval <= 0 {
interval = DefaultInterval
}
return &Enforcer{
source: source,
limits: limits,
revoker: revoker,
interval: interval,
now: time.Now,
}
}
// Start runs the evaluation loop until ctx is cancelled. Non-blocking.
func (e *Enforcer) Start(ctx context.Context) {
if e == nil || e.source == nil || e.limits == nil || e.revoker == nil {
return
}
go func() {
ticker := time.NewTicker(e.interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
e.evaluate(ctx)
}
}
}()
}
// evaluate runs one pass: snapshot → per-user over-cap check → revoke victims.
// Exported behavior is covered by EvaluateOnce for tests.
func (e *Enforcer) evaluate(ctx context.Context) {
if err := e.EvaluateOnce(ctx); err != nil {
slog.Debug("stream enforcer evaluate failed", "error", err)
}
}
// EvaluateOnce performs a single enforcement pass and returns the number of
// sessions revoked. Deterministic and side-effect-scoped for testing.
func (e *Enforcer) EvaluateOnce(ctx context.Context) error {
snap, err := e.source.Snapshot(ctx)
if err != nil {
return err
}
for userID, streams := range snap.ByUser() {
if userID <= 0 || len(streams) == 0 {
// user 0 == records with no resolved owner; never enforce against it.
continue
}
limit, err := e.limits(ctx, userID)
if err != nil {
// Fail open: a limit-lookup error must never kill legitimate streams.
slog.Debug("stream enforcer: limit lookup failed; skipping user",
"user_id", userID, "error", err)
continue
}
if limit <= 0 || len(streams) <= limit {
continue
}
for _, victim := range e.selectVictims(streams, limit) {
if err := e.revoker.RevokeSessionFor(ctx, victim.SessionID, "over_concurrent_stream_limit", revocationTTL); err != nil {
slog.Warn("stream enforcer: revoke failed",
"user_id", userID, "session_id", victim.SessionID, "error", err)
continue
}
slog.Info("stream enforcer: revoked over-cap session",
"user_id", userID, "session_id", victim.SessionID,
"limit", limit, "live", len(streams))
}
}
return nil
}
// selectVictims returns the streams beyond the limit, keeping the `limit`
// MOST-RECENTLY-SERVED sessions and trimming the rest. Ordering by real serve
// activity (LastServedAt), not StartedAt, matters: after a network blip a client
// reconnects with a new session while the old ghost lingers in the monitor for
// up to its TTL. Keeping the freshly-served sessions means the live reconnect
// survives and the stale ghost is the one trimmed (and it would have aged out
// anyway). Falls back to StartedAt, then session id, for deterministic ties.
func (e *Enforcer) selectVictims(streams []streammonitor.LiveStream, limit int) []streammonitor.LiveStream {
ordered := make([]streammonitor.LiveStream, len(streams))
copy(ordered, streams)
// Most-recently-served first.
sort.SliceStable(ordered, func(i, j int) bool {
if !ordered[i].LastServedAt.Equal(ordered[j].LastServedAt) {
return ordered[i].LastServedAt.After(ordered[j].LastServedAt)
}
if !ordered[i].StartedAt.Equal(ordered[j].StartedAt) {
return ordered[i].StartedAt.After(ordered[j].StartedAt)
}
return ordered[i].SessionID < ordered[j].SessionID
})
if limit >= len(ordered) {
return nil
}
// Keep the first `limit` (freshest); revoke the rest (stalest).
return ordered[limit:]
}
+138
View File
@@ -0,0 +1,138 @@
package streamenforcer
import (
"context"
"errors"
"sort"
"testing"
"time"
"github.com/Silo-Server/silo-server/internal/streammonitor"
)
type fakeSource struct {
snap streammonitor.Snapshot
err error
}
func (f fakeSource) Snapshot(context.Context) (streammonitor.Snapshot, error) {
return f.snap, f.err
}
type fakeRevoker struct {
revoked []string
err error
}
func (f *fakeRevoker) RevokeSessionFor(_ context.Context, sessionID, _ string, _ time.Duration) error {
if f.err != nil {
return f.err
}
f.revoked = append(f.revoked, sessionID)
return nil
}
// stream builds a LiveStream whose StartedAt and LastServedAt are both ts.
func stream(id string, user int, ts time.Time) streammonitor.LiveStream {
return streammonitor.LiveStream{SessionID: id, UserID: user, StartedAt: ts, LastServedAt: ts}
}
// served builds a LiveStream started at `started` but last served at `lastServed`
// (to model a ghost: old activity) — StartedAt independent from real liveness.
func served(id string, user int, started, lastServed time.Time) streammonitor.LiveStream {
return streammonitor.LiveStream{SessionID: id, UserID: user, StartedAt: started, LastServedAt: lastServed}
}
func TestEvaluateOnce(t *testing.T) {
base := time.Date(2026, 7, 4, 12, 0, 0, 0, time.UTC)
limit3 := func(context.Context, int) (int, error) { return 3, nil }
tests := []struct {
name string
streams []streammonitor.LiveStream
limits LimitFunc
wantRevoke []string
}{
{
name: "under cap: no revokes",
streams: []streammonitor.LiveStream{stream("a", 1, base), stream("b", 1, base.Add(time.Minute))},
limits: limit3,
wantRevoke: nil,
},
{
name: "over cap: least-recently-served trimmed, freshest kept",
streams: []streammonitor.LiveStream{
stream("stale", 1, base),
stream("old", 1, base.Add(1*time.Minute)),
stream("keep1", 1, base.Add(2*time.Minute)),
stream("keep2", 1, base.Add(3*time.Minute)),
stream("keep3", 1, base.Add(4*time.Minute)),
},
limits: limit3,
wantRevoke: []string{"stale", "old"},
},
{
name: "ghost trimmed, fresh reconnect survives",
// Two fresh (served just now) and two ghosts (served long ago), limit 2.
streams: []streammonitor.LiveStream{
served("ghost1", 1, base, base),
served("ghost2", 1, base.Add(time.Second), base.Add(time.Second)),
served("fresh1", 1, base.Add(10*time.Minute), base.Add(10*time.Minute)),
served("fresh2", 1, base.Add(10*time.Minute), base.Add(11*time.Minute)),
},
limits: func(context.Context, int) (int, error) { return 2, nil },
wantRevoke: []string{"ghost1", "ghost2"},
},
{
name: "limit lookup error: fail open",
streams: []streammonitor.LiveStream{stream("a", 1, base), stream("b", 1, base), stream("c", 1, base), stream("d", 1, base)},
limits: func(context.Context, int) (int, error) { return 0, errors.New("db down") },
wantRevoke: nil,
},
{
name: "unlimited (limit<=0): no revokes",
streams: []streammonitor.LiveStream{stream("a", 1, base), stream("b", 1, base), stream("c", 1, base)},
limits: func(context.Context, int) (int, error) { return 0, nil },
wantRevoke: nil,
},
{
name: "user 0 never enforced",
streams: []streammonitor.LiveStream{stream("a", 0, base), stream("b", 0, base), stream("c", 0, base), stream("d", 0, base)},
limits: limit3,
wantRevoke: nil,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
rev := &fakeRevoker{}
e := New(fakeSource{snap: streammonitor.Snapshot{Streams: tc.streams}}, tc.limits, rev, 0)
if err := e.EvaluateOnce(context.Background()); err != nil {
t.Fatalf("EvaluateOnce: %v", err)
}
got := append([]string{}, rev.revoked...)
sort.Strings(got)
want := append([]string{}, tc.wantRevoke...)
sort.Strings(want)
if len(got) != len(want) {
t.Fatalf("revoked = %v, want %v", got, want)
}
for i := range got {
if got[i] != want[i] {
t.Fatalf("revoked = %v, want %v", got, want)
}
}
})
}
}
func TestEvaluateOnceSourceError(t *testing.T) {
rev := &fakeRevoker{}
e := New(fakeSource{err: errors.New("scan failed")}, func(context.Context, int) (int, error) { return 1, nil }, rev, 0)
if err := e.EvaluateOnce(context.Background()); err == nil {
t.Fatal("expected error from source")
}
if len(rev.revoked) != 0 {
t.Fatalf("expected no revokes on source error, got %v", rev.revoked)
}
}
+107
View File
@@ -0,0 +1,107 @@
package streamrevoke
import (
"context"
"fmt"
"time"
"github.com/jackc/pgx/v5/pgxpool"
)
// permanentExpiry is the far-future sentinel written when a Revocation has a
// zero-value ExpiresAt. The hot path treats a zero ExpiresAt as "never expires"
// (a permanent kill), but the DB column is NOT NULL and Prune/ListActive compare
// expires_at <= now(): a literal zero time (0001-01-01) would be excluded by
// ListActive and deleted by the very next Prune, silently evaporating a
// permanent kill. Writing a year-2999 sentinel preserves the intent durably.
var permanentExpiry = time.Date(2999, 1, 1, 0, 0, 0, 0, time.UTC)
// PostgresDurableStore is the Postgres-backed DurableStore: a durable mirror of
// the kill list so revocations survive a Redis flush or a server restart. It is
// never on the hot path — Store consults it only on write (Upsert), on
// warm/reconcile (ListActive), and on trim (Prune).
//
// Rows are keyed by (kind, id) so re-revoking the same session/user (the async
// over-cap enforcer does this every pass) UPSERTs the same row rather than
// accumulating duplicates; physical growth is reclaimed by Prune.
type PostgresDurableStore struct {
pool *pgxpool.Pool
}
// NewPostgresDurableStore builds a DurableStore from a pgx pool. It returns a
// nil DurableStore interface when pool is nil so callers can pass the result
// straight into Options.Durable and a Redis-less/DB-less mode degrades to a
// true nil interface (avoiding the "non-nil interface wrapping a nil pointer"
// trap that would make Store.durable != nil erroneously true).
func NewPostgresDurableStore(pool *pgxpool.Pool) DurableStore {
if pool == nil {
return nil
}
return &PostgresDurableStore{pool: pool}
}
// Upsert writes or refreshes a revocation, keyed by (kind, id).
func (s *PostgresDurableStore) Upsert(ctx context.Context, r Revocation) error {
// A zero ExpiresAt means "permanent" on the hot path; persist it as a
// far-future sentinel so Prune/ListActive don't immediately reap the row.
expiresAt := r.ExpiresAt
if expiresAt.IsZero() {
expiresAt = permanentExpiry
}
// Expiry is monotonic: a re-revoke never shortens an existing longer kill
// (GREATEST). The async over-cap enforcer re-revokes with a short 5m TTL, and
// without this it would shrink an admin's 24h kill on the same session key and
// reopen the restart-resurrection window. reason/revoked_at follow whichever
// expiry wins so the persisted row stays coherent. Mirrors applyLocal.
_, err := s.pool.Exec(ctx, `
INSERT INTO stream_revocations (kind, id, reason, revoked_at, expires_at)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (kind, id) DO UPDATE SET
reason = CASE WHEN EXCLUDED.expires_at >= stream_revocations.expires_at
THEN EXCLUDED.reason ELSE stream_revocations.reason END,
revoked_at = CASE WHEN EXCLUDED.expires_at >= stream_revocations.expires_at
THEN EXCLUDED.revoked_at ELSE stream_revocations.revoked_at END,
expires_at = GREATEST(stream_revocations.expires_at, EXCLUDED.expires_at)`,
string(r.Kind), r.ID, r.Reason, r.RevokedAt, expiresAt)
if err != nil {
return fmt.Errorf("streamrevoke upsert: %w", err)
}
return nil
}
// ListActive returns every revocation not yet expired, for warming the hot-path
// cache on startup and re-warming after a Redis flush.
func (s *PostgresDurableStore) ListActive(ctx context.Context) ([]Revocation, error) {
rows, err := s.pool.Query(ctx, `
SELECT kind, id, reason, revoked_at, expires_at
FROM stream_revocations
WHERE expires_at > now()`)
if err != nil {
return nil, fmt.Errorf("streamrevoke list active: %w", err)
}
defer rows.Close()
var out []Revocation
for rows.Next() {
var r Revocation
var kind string
if err := rows.Scan(&kind, &r.ID, &r.Reason, &r.RevokedAt, &r.ExpiresAt); err != nil {
return nil, fmt.Errorf("streamrevoke scan: %w", err)
}
r.Kind = Kind(kind)
out = append(out, r)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("streamrevoke list active rows: %w", err)
}
return out, nil
}
// Prune physically deletes expired rows so the table does not grow unbounded as
// the async enforcer re-revokes across passes.
func (s *PostgresDurableStore) Prune(ctx context.Context) error {
if _, err := s.pool.Exec(ctx, `DELETE FROM stream_revocations WHERE expires_at <= now()`); err != nil {
return fmt.Errorf("streamrevoke prune: %w", err)
}
return nil
}
+586
View File
@@ -0,0 +1,586 @@
// Package streamrevoke implements a stream-revocation "kill switch": a small
// set of revoked stream sessions/users that edge nodes consult on the hot path
// via an in-memory cache, backed by Redis for multi-node propagation with an
// optional durable Postgres mirror.
//
// The hot path (IsRevoked) is a pure in-memory read with no I/O. Revocations
// are propagated to other nodes over Redis pub/sub for immediate application
// and mirrored to Redis keys (with a TTL) so late-joining or restarting nodes
// can reconcile via a periodic SCAN. A durable mirror, when configured, lets
// kills survive a Redis flush.
package streamrevoke
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"net/http"
"strconv"
"sync"
"time"
"github.com/Silo-Server/silo-server/internal/cache"
"github.com/redis/go-redis/v9"
)
const (
// keyPrefix is the Redis key namespace for revocation mirror keys, e.g.
// silo:revoked:sess:{id} and silo:revoked:user:{id}.
keyPrefix = "silo:revoked:"
// scanPattern matches every revocation mirror key for cache warming and
// periodic reconciliation.
scanPattern = keyPrefix + "*"
// EventStreamRevoked is the cache.Event.Type published on cache.ChannelAdmin
// when a revocation is created so other nodes apply it immediately.
EventStreamRevoked = "stream_revoked"
defaultPollInterval = 60 * time.Second
// defaultTTL is intentionally >= the playback recipe-card lifetime
// (playback.MaxTokenTTL, 24h): a kill must never expire before the session it
// kills can be reconstructed, or PR #174's restart-resilient playback could
// rebuild and re-serve a stream whose revocation had already lapsed. This is
// an invariant, stated in words to avoid importing the playback package here
// (which would create an import cycle); keep the two values coupled.
defaultTTL = 24 * time.Hour
)
// Kind identifies what a revocation targets.
type Kind string
const (
// KindSession revokes a single stream session id.
KindSession Kind = "sess"
// KindUser revokes every stream belonging to a user id.
KindUser Kind = "user"
)
// Key uniquely identifies a revocation in the in-memory cache and Redis.
type Key struct {
Kind Kind
ID string
}
// Revocation is a single kill-switch entry.
type Revocation struct {
Kind Kind `json:"kind"`
ID string `json:"id"`
Reason string `json:"reason,omitempty"`
RevokedAt time.Time `json:"revoked_at"`
ExpiresAt time.Time `json:"expires_at"`
}
// key returns the in-memory cache key for this revocation.
func (r Revocation) key() Key {
return Key{Kind: r.Kind, ID: r.ID}
}
// expired reports whether the revocation is no longer in force at now.
func (r Revocation) expired(now time.Time) bool {
return !r.ExpiresAt.IsZero() && !now.Before(r.ExpiresAt)
}
// DurableStore is an optional Postgres-backed mirror so kills survive a Redis
// flush. A concrete implementation lives elsewhere; this package only consumes
// the interface when one is provided.
type DurableStore interface {
Upsert(ctx context.Context, r Revocation) error
ListActive(ctx context.Context) ([]Revocation, error)
Prune(ctx context.Context) error
}
// Options configures a Store.
type Options struct {
Redis *redis.Client // nil => memory-only (integrated single-node)
Bus cache.EventBus // nil => no push propagation
Durable DurableStore // nil => no durable mirror
PollInterval time.Duration // default 60s
DefaultTTL time.Duration // default 24h
}
// Store holds the in-memory revocation cache and its propagation plumbing.
type Store struct {
rdb *redis.Client
bus cache.EventBus
durable DurableStore
pollInterval time.Duration
defaultTTL time.Duration
mu sync.RWMutex
items map[Key]Revocation
}
// New builds a Store, applying defaults for any unset Options.
func New(opts Options) *Store {
if opts.PollInterval <= 0 {
opts.PollInterval = defaultPollInterval
}
if opts.DefaultTTL <= 0 {
opts.DefaultTTL = defaultTTL
}
return &Store{
rdb: opts.Redis,
bus: opts.Bus,
durable: opts.Durable,
pollInterval: opts.PollInterval,
defaultTTL: opts.DefaultTTL,
items: make(map[Key]Revocation),
}
}
// redisKey returns the Redis mirror key for a revocation key.
func redisKey(k Key) string {
return keyPrefix + string(k.Kind) + ":" + k.ID
}
// userKey returns the cache key for a numeric user id.
func userKey(userID int) Key {
return Key{Kind: KindUser, ID: strconv.Itoa(userID)}
}
// Refuse is the single shared enforcement point for every serve surface (edge
// proxy, transcode node, native api, jellycompat): if the session or user is
// revoked it writes a 403 and returns true, and the caller must stop serving.
// Centralizing the check + response here keeps one section per concern instead
// of duplicating "IsRevoked → 403" at each surface. It does NOT hang up an
// in-flight connection — long-pour paths pair this with a connection-cut helper.
// startedAt is when the request's stream credential was issued (see IsRevoked).
func (s *Store) Refuse(w http.ResponseWriter, sessionID string, userID int, startedAt time.Time) bool {
if s == nil || !s.IsRevoked(sessionID, userID, startedAt) {
return false
}
http.Error(w, "stream revoked", http.StatusForbidden)
return true
}
// WatchAndCut watches a long-lived pour (a single long-GET direct-play/remux) and
// once the session/user is revoked, forces the in-flight write to fail via
// SetWriteDeadline — hanging up the socket even though the 24h token is still
// valid (cutting an open connection is a socket action, not a token revocation).
// It checks immediately on entry (so a pour that began the instant before the
// kill is cut without waiting a tick) and then every 5s. Returns a stop func the
// caller defers when the request finishes normally. This is the shared in-flight
// cut used by every long-pour serve surface (edge proxy, native api,
// jellycompat), so the cut logic lives in one place. HLS/transcode paths don't
// need it — per-segment Refuse stops them within one segment.
//
// Best-effort: if the ResponseWriter chain doesn't support write deadlines the
// deadline set is a no-op and the stream still stops on its next request via
// Refuse. Never wraps the writer, so it does not disable sendfile.
// startedAt follows IsRevoked's contract; a pour in flight when a user kill
// lands always predates that kill, so passing the request's credential/entry
// time makes mid-pour user kills cut correctly on every surface.
func (s *Store) WatchAndCut(w http.ResponseWriter, sessionID string, userID int, startedAt time.Time) func() {
if s == nil {
return func() {}
}
cut := func() { _ = http.NewResponseController(w).SetWriteDeadline(time.Now()) }
if s.IsRevoked(sessionID, userID, startedAt) {
cut()
return func() {}
}
done := make(chan struct{})
go func() {
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-done:
return
case <-ticker.C:
if s.IsRevoked(sessionID, userID, startedAt) {
cut()
return
}
}
}
}()
var once sync.Once
return func() { once.Do(func() { close(done) }) }
}
// IsRevoked is the HOT PATH: a pure in-memory cache read with no I/O. It
// returns true if the session id is currently revoked, or if the user id is
// revoked AND the stream's credential predates that user revocation.
//
// startedAt is when the request's stream credential was issued: the stream
// token's iat on token-bearing surfaces (edge proxy, transcode node, native
// ?st=), or the request entry time on freshly-authenticated surfaces (native
// session auth, jellycompat login). A user revocation is a CUTOFF, not a ban:
// it kills streams whose credential predates it, while a stream authorized
// after it (which required passing auth that the revocation just reset) plays
// normally. Without the cutoff, the OnUserSessionsRevoked hook — fired by any
// admin edit of password/role/enabled/permissions/quality — would 403 the
// user's playback for the full 24h TTL even after they re-authenticate, with
// no unrevoke path (expiry is deliberately monotonic). A zero startedAt never
// matches a user revocation (fail open, matching the enforcer's "never kill on
// uncertainty"); session revocations are exact-id kills and ignore startedAt.
func (s *Store) IsRevoked(sessionID string, userID int, startedAt time.Time) bool {
now := time.Now()
s.mu.RLock()
defer s.mu.RUnlock()
if sessionID != "" {
if r, ok := s.items[Key{Kind: KindSession, ID: sessionID}]; ok && !r.expired(now) {
return true
}
}
// userID <= 0 is the "no resolved owner" sentinel (e.g. the transcode node's
// session-only check). Never match it against a KindUser entry: a stray
// user:"0" revocation must not read as "every ownerless request is revoked".
if userID > 0 {
if r, ok := s.items[userKey(userID)]; ok && !r.expired(now) &&
!startedAt.IsZero() && startedAt.Before(r.RevokedAt) {
return true
}
}
return false
}
// effectiveExpiry returns the instant a revocation lapses, treating a zero
// ExpiresAt as "never" (far future) so monotonic comparisons order a permanent
// kill above any bounded one.
func (r Revocation) effectiveExpiry() time.Time {
if r.ExpiresAt.IsZero() {
return permanentExpiry
}
return r.ExpiresAt
}
// applyLocal inserts or refreshes a revocation in the in-memory cache, keeping
// whichever copy expires LATER. Expiry is monotonic: a re-revoke with a shorter
// TTL must never shorten a longer-lived kill. This matters because the async
// over-cap enforcer re-revokes with a short self-healing TTL (5m); without this
// guard it would silently shrink an admin's 24h RevokeSession on the same
// session key and reopen the restart-resurrection window PR #174 exposes. The
// same rule lets a mid-window Redis reconcile that reads a stale longer entry
// not fight a freshly-applied one. The caller must not hold s.mu.
func (s *Store) applyLocal(r Revocation) {
s.mu.Lock()
if existing, ok := s.items[r.key()]; !ok || !existing.effectiveExpiry().After(r.effectiveExpiry()) {
s.items[r.key()] = r
}
s.mu.Unlock()
}
// effective returns the currently-stored revocation for a key — the
// monotonically-merged copy applyLocal kept (whichever expires later). Zero
// value if absent.
func (s *Store) effective(k Key) Revocation {
s.mu.RLock()
defer s.mu.RUnlock()
return s.items[k]
}
// Revoke adds a revocation to the local cache, writes it to Redis (with a TTL
// of until-now) when configured, mirrors it to the durable store when
// configured, and publishes a pub/sub event so other nodes apply it
// immediately.
func (s *Store) Revoke(ctx context.Context, key Key, reason string, until time.Time) error {
now := time.Now()
r := Revocation{
Kind: key.Kind,
ID: key.ID,
Reason: reason,
RevokedAt: now,
ExpiresAt: until,
}
// Local cache first: the hot path must reflect the kill immediately even if
// downstream propagation fails.
s.applyLocal(r)
// Propagation must not die with the caller: an admin terminate rides its
// HTTP request's context, and an abort after applyLocal would strand the
// kill in this process's memory — never reaching Redis (edges), pub/sub, or
// the durable mirror (lost on restart).
ctx = context.WithoutCancel(ctx)
// Durable mirror (best-effort; failures are non-fatal).
if s.durable != nil {
if err := s.durable.Upsert(ctx, r); err != nil {
slog.Warn("streamrevoke durable upsert failed", "error", err, "kind", r.Kind, "id", r.ID)
}
}
// Redis mirror for multi-node propagation to edges. Mirror the merged copy
// applyLocal kept (monotonic): a short over-cap re-revoke must not shorten a
// longer admin kill's Redis TTL, matching applyLocal and the durable upsert's
// GREATEST. Edges that warm from Redis in that window must see the later kill.
s.mirrorToRedis(ctx, s.effective(r.key()))
// Pub/sub push so already-connected nodes apply the kill immediately (the
// warm loops deliberately skip this — SCAN reconcile handles late-joiners).
if s.bus != nil {
data, err := json.Marshal(r)
if err != nil {
slog.Warn("streamrevoke marshal failed", "error", err, "kind", r.Kind, "id", r.ID)
return err
}
evt := cache.Event{Type: EventStreamRevoked, Payload: string(data)}
if err := s.bus.Publish(ctx, cache.ChannelAdmin, evt); err != nil {
slog.Warn("streamrevoke publish failed", "error", err, "kind", r.Kind, "id", r.ID)
}
}
return nil
}
// mirrorToRedis writes (or deletes) the Redis mirror key for a revocation with a
// TTL of the time remaining until r.ExpiresAt. It is the single shared
// Redis-arming path, used both by Revoke and by the durable warm loops: edge
// nodes (proxy/transcode) have no durable store and learn kills ONLY from
// Redis, so re-arming Redis from the durable source of truth after a Redis flush
// + central restart is what lets edges reconverge via their SCAN reconcile.
// Best-effort: failures are logged, never returned. No-op when Redis is absent.
func (s *Store) mirrorToRedis(ctx context.Context, r Revocation) {
if s.rdb == nil {
return
}
data, err := json.Marshal(r)
if err != nil {
slog.Warn("streamrevoke marshal failed", "error", err, "kind", r.Kind, "id", r.ID)
return
}
if r.ExpiresAt.IsZero() {
// Permanent kill: set without a TTL, matching applyLocal/effectiveExpiry
// treating a zero ExpiresAt as "never". A TTL of time.Until(zero) would be
// hugely negative and drop the key, so late-joining edges would miss it.
if err := s.rdb.Set(ctx, redisKey(r.key()), data, 0).Err(); err != nil {
slog.Warn("streamrevoke redis set failed", "error", err, "kind", r.Kind, "id", r.ID)
}
return
}
ttl := time.Until(r.ExpiresAt)
if ttl <= 0 {
// Already expired: nothing to mirror in Redis.
if err := s.rdb.Del(ctx, redisKey(r.key())).Err(); err != nil {
slog.Debug("streamrevoke redis del failed", "error", err, "kind", r.Kind, "id", r.ID)
}
return
}
if err := s.rdb.Set(ctx, redisKey(r.key()), data, ttl).Err(); err != nil {
slog.Warn("streamrevoke redis set failed", "error", err, "kind", r.Kind, "id", r.ID)
}
}
// RevokeSession revokes a single session id for DefaultTTL.
func (s *Store) RevokeSession(ctx context.Context, sessionID, reason string) error {
return s.Revoke(ctx, Key{Kind: KindSession, ID: sessionID}, reason, time.Now().Add(s.defaultTTL))
}
// RevokeSessionFor revokes a single session id for a caller-chosen TTL. The
// async enforcer uses a short TTL so a transient over-count (e.g. a ghost
// session lingering next to a fresh reconnect) self-heals; a persistent abuser
// is simply re-revoked on the next evaluation pass. ttl <= 0 falls back to
// DefaultTTL.
func (s *Store) RevokeSessionFor(ctx context.Context, sessionID, reason string, ttl time.Duration) error {
if ttl <= 0 {
ttl = s.defaultTTL
}
return s.Revoke(ctx, Key{Kind: KindSession, ID: sessionID}, reason, time.Now().Add(ttl))
}
// RevokeUser revokes the user's streams for DefaultTTL. It is a CUTOFF, not a
// ban: enforcement (IsRevoked) only matches streams whose credential was issued
// before this revocation, so playback the user starts after re-authenticating
// is unaffected. This is what makes it safe for OnUserSessionsRevoked to call
// on every auth-session revocation (admin edits included), while still cutting
// every stream that rode a pre-revocation credential.
func (s *Store) RevokeUser(ctx context.Context, userID int, reason string) error {
// IsRevoked never matches userID <= 0 against a KindUser entry, so a kill for
// such an id would be silently ineffective. Refuse it rather than report a
// success that has no teeth.
if userID <= 0 {
return fmt.Errorf("streamrevoke: invalid userID %d", userID)
}
return s.Revoke(ctx, userKey(userID), reason, time.Now().Add(s.defaultTTL))
}
// List returns the currently-active revocations, pruning expired entries on
// read.
func (s *Store) List() []Revocation {
now := time.Now()
s.mu.Lock()
defer s.mu.Unlock()
out := make([]Revocation, 0, len(s.items))
for k, r := range s.items {
if r.expired(now) {
delete(s.items, k)
continue
}
out = append(out, r)
}
return out
}
// pruneExpired removes expired entries from the in-memory cache.
func (s *Store) pruneExpired() {
now := time.Now()
s.mu.Lock()
for k, r := range s.items {
if r.expired(now) {
delete(s.items, k)
}
}
s.mu.Unlock()
}
// StartSync warms the cache (durable store then a Redis SCAN of silo:revoked:*)
// and subscribes to cache.ChannelAdmin to apply EventStreamRevoked events. The
// initial warm and subscribe run SYNCHRONOUSLY (inline) before StartSync
// returns, so kills recorded durably or in Redis are already enforced before
// the first stream is served. Only the PollInterval reconcile/prune loop runs
// in a spawned goroutine. With a nil Redis client the Store is memory-only and
// only the local prune (plus durable maintenance, if configured) runs on the
// tick.
func (s *Store) StartSync(ctx context.Context) {
// Warm from the durable mirror first so kills survive a Redis flush. Each
// non-expired entry is also re-armed into Redis (mirrorToRedis no-ops when
// Redis is absent) so edge nodes — which learn kills only from Redis —
// reconverge via their SCAN reconcile after a Redis flush + central restart.
if s.durable != nil {
if revs, err := s.warmFromDurable(ctx); err != nil {
slog.Warn("streamrevoke durable warm failed", "error", err)
} else {
now := time.Now()
for _, r := range revs {
if !r.expired(now) {
s.applyLocal(r)
s.mirrorToRedis(ctx, r)
}
}
}
}
// Warm from Redis and subscribe for push updates.
if s.rdb != nil {
s.reconcileFromRedis(ctx)
}
if s.bus != nil {
if err := s.bus.Subscribe(ctx, cache.ChannelAdmin, s.handleEvent); err != nil {
slog.Warn("streamrevoke subscribe failed", "error", err)
}
}
go func() {
ticker := time.NewTicker(s.pollInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
s.maintain(ctx)
}
}
}()
}
// warmFromDurable loads the active revocations for the boot-time cache warm,
// retrying a bounded number of times with short backoff. A single transient DB
// error at startup (with Redis also flushed) would otherwise leave the kill
// list empty until the first poll tick 60s later; the retry closes that window.
// It is strictly bounded and honors ctx cancellation, so startup never blocks
// indefinitely. Only the boot warm needs this — the recurring maintain tick is
// its own backstop.
func (s *Store) warmFromDurable(ctx context.Context) ([]Revocation, error) {
backoffs := []time.Duration{200 * time.Millisecond, 400 * time.Millisecond}
var revs []Revocation
var err error
for attempt := 0; ; attempt++ {
if revs, err = s.durable.ListActive(ctx); err == nil {
return revs, nil
}
if attempt >= len(backoffs) {
return nil, err
}
select {
case <-ctx.Done():
return nil, err
case <-time.After(backoffs[attempt]):
}
}
}
// maintain runs one reconcile/prune pass: it is the body of the poll tick,
// factored out so tests can exercise it directly without a live ticker. All
// steps are best-effort — a failure in one is logged and does not block the
// others or the next tick.
func (s *Store) maintain(ctx context.Context) {
if s.rdb != nil {
s.reconcileFromRedis(ctx)
}
// Durable maintenance: re-warm from the durable mirror so a mid-life Redis
// flush self-heals, re-arming Redis for edge nodes too, then physically
// reclaim expired rows.
if s.durable != nil {
if revs, err := s.durable.ListActive(ctx); err != nil {
slog.Warn("streamrevoke durable reconcile failed", "error", err)
} else {
now := time.Now()
for _, r := range revs {
if !r.expired(now) {
s.applyLocal(r)
s.mirrorToRedis(ctx, r)
}
}
}
if err := s.durable.Prune(ctx); err != nil {
slog.Warn("streamrevoke durable prune failed", "error", err)
}
}
s.pruneExpired()
}
// handleEvent applies an EventStreamRevoked event to the local cache. Other
// event types on the channel are ignored.
func (s *Store) handleEvent(evt cache.Event) {
if evt.Type != EventStreamRevoked {
return
}
var r Revocation
if err := json.Unmarshal([]byte(evt.Payload), &r); err != nil {
slog.Debug("streamrevoke event unmarshal failed", "error", err)
return
}
if r.expired(time.Now()) {
return
}
s.applyLocal(r)
}
// reconcileFromRedis SCANs the revocation namespace and applies every live
// entry to the local cache. Redis is authoritative for entries it holds; TTL
// expiry there is mirrored by prune-on-read locally.
func (s *Store) reconcileFromRedis(ctx context.Context) {
var cursor uint64
now := time.Now()
for {
keys, next, err := s.rdb.Scan(ctx, cursor, scanPattern, 256).Result()
if err != nil {
slog.Warn("streamrevoke redis scan failed", "error", err)
return
}
for _, k := range keys {
val, err := s.rdb.Get(ctx, k).Result()
if err != nil {
// Key may have expired between SCAN and GET; skip it.
continue
}
var r Revocation
if err := json.Unmarshal([]byte(val), &r); err != nil {
slog.Debug("streamrevoke redis unmarshal failed", "error", err, "key", k)
continue
}
if !r.expired(now) {
s.applyLocal(r)
}
}
cursor = next
if cursor == 0 {
break
}
}
}
+471
View File
@@ -0,0 +1,471 @@
package streamrevoke
import (
"context"
"errors"
"sync"
"testing"
"time"
)
// errFakeDurable is the error fakeDurable.ListActive returns while failNext > 0.
var errFakeDurable = errors.New("fake durable failure")
// newMemStore returns a memory-only Store (no Redis, no bus, no durable mirror).
func newMemStore() *Store {
return New(Options{})
}
// fakeDurable is an in-memory DurableStore double for exercising the durable
// wiring (Upsert on revoke, warm on start, Prune on the poll tick) without a
// live Postgres. StartSync's warm is synchronous, but the poll goroutine calls
// ListActive/Prune concurrently with the test's own reads, so every field is
// mutex-guarded. failNext, when > 0, makes ListActive return an error that many
// times before succeeding, to exercise the bounded boot-warm retry.
type fakeDurable struct {
mu sync.Mutex
rows map[Key]Revocation
upserts int
prunes int
failNext int
lists int
}
func newFakeDurable() *fakeDurable {
return &fakeDurable{rows: make(map[Key]Revocation)}
}
func (f *fakeDurable) Upsert(_ context.Context, r Revocation) error {
f.mu.Lock()
defer f.mu.Unlock()
f.rows[r.key()] = r
f.upserts++
return nil
}
func (f *fakeDurable) ListActive(_ context.Context) ([]Revocation, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.lists++
if f.failNext > 0 {
f.failNext--
return nil, errFakeDurable
}
now := time.Now()
out := make([]Revocation, 0, len(f.rows))
for _, r := range f.rows {
if !r.expired(now) {
out = append(out, r)
}
}
return out, nil
}
func (f *fakeDurable) Prune(_ context.Context) error {
f.mu.Lock()
defer f.mu.Unlock()
now := time.Now()
for k, r := range f.rows {
if r.expired(now) {
delete(f.rows, k)
}
}
f.prunes++
return nil
}
func (f *fakeDurable) upsertCount() int {
f.mu.Lock()
defer f.mu.Unlock()
return f.upserts
}
func (f *fakeDurable) pruneCount() int {
f.mu.Lock()
defer f.mu.Unlock()
return f.prunes
}
func (f *fakeDurable) listCount() int {
f.mu.Lock()
defer f.mu.Unlock()
return f.lists
}
// TestRevokeMirrorsToDurable asserts Revoke writes through to the durable store.
func TestRevokeMirrorsToDurable(t *testing.T) {
ctx := context.Background()
fake := newFakeDurable()
s := New(Options{Durable: fake})
if err := s.RevokeSession(ctx, "sess-1", "abuse"); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
if got := fake.upsertCount(); got != 1 {
t.Fatalf("durable upserts = %d, want 1", got)
}
// Re-revoking the same session upserts the same row (bounded growth).
if err := s.RevokeSession(ctx, "sess-1", "abuse again"); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
fake.mu.Lock()
rows := len(fake.rows)
fake.mu.Unlock()
if rows != 1 {
t.Fatalf("durable rows after re-revoke = %d, want 1", rows)
}
}
// TestStartSyncWarmsFromDurable simulates a restart: a fresh Store must
// repopulate its hot-path map from the durable mirror, and skip expired rows.
func TestStartSyncWarmsFromDurable(t *testing.T) {
// Cancel the poll goroutine StartSync spawns so it does not leak past the test.
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
fake := newFakeDurable()
// Seed the durable store as if a previous process had written these.
active := Revocation{Kind: KindSession, ID: "sess-live", Reason: "x", RevokedAt: time.Now(), ExpiresAt: time.Now().Add(time.Hour)}
expired := Revocation{Kind: KindUser, ID: "9", Reason: "x", RevokedAt: time.Now().Add(-2 * time.Hour), ExpiresAt: time.Now().Add(-time.Hour)}
_ = fake.Upsert(ctx, active)
_ = fake.Upsert(ctx, expired)
s := New(Options{Durable: fake})
s.StartSync(ctx)
if !s.IsRevoked("sess-live", 0, time.Time{}) {
t.Fatalf("expected sess-live to be revoked after warm from durable")
}
// startedAt predates the (expired) user revocation, so only expiry decides.
if s.IsRevoked("whatever", 9, time.Now().Add(-3*time.Hour)) {
t.Fatalf("expected expired user 9 not to be warmed into the cache")
}
}
// TestWarmFromDurableRetries proves the bounded boot-warm retry (Fix 2): a
// transient ListActive failure at startup is retried rather than leaving the
// kill list empty until the first poll tick.
func TestWarmFromDurableRetries(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
fake := newFakeDurable()
live := Revocation{Kind: KindSession, ID: "sess-retry", Reason: "x", RevokedAt: time.Now(), ExpiresAt: time.Now().Add(time.Hour)}
_ = fake.Upsert(ctx, live)
// Fail the first two ListActive calls; the third (within the 3-attempt
// bound) succeeds.
fake.mu.Lock()
fake.failNext = 2
fake.lists = 0
fake.mu.Unlock()
s := New(Options{Durable: fake})
s.StartSync(ctx)
if !s.IsRevoked("sess-retry", 0, time.Time{}) {
t.Fatalf("expected sess-retry warmed into cache after retried ListActive")
}
if got := fake.listCount(); got != 3 {
t.Fatalf("ListActive calls = %d, want 3 (two failures + one success)", got)
}
}
// TestWarmFromDurableGivesUpBounded proves the retry is bounded: when every
// attempt fails, warmFromDurable stops after the fixed attempt budget instead
// of blocking startup.
func TestWarmFromDurableGivesUpBounded(t *testing.T) {
fake := newFakeDurable()
fake.mu.Lock()
fake.failNext = 100 // always fail
fake.mu.Unlock()
s := New(Options{Durable: fake})
if _, err := s.warmFromDurable(context.Background()); err == nil {
t.Fatalf("expected warmFromDurable to return an error when all attempts fail")
}
if got := fake.listCount(); got != 3 {
t.Fatalf("ListActive attempts = %d, want 3 (bounded)", got)
}
}
// TestMaintainPrunesAndReWarmsFromDurable asserts the poll-tick body calls the
// durable Prune and re-warms live rows into the hot-path map (self-healing a
// Redis flush that cleared the local cache).
func TestMaintainPrunesAndReWarmsFromDurable(t *testing.T) {
ctx := context.Background()
fake := newFakeDurable()
live := Revocation{Kind: KindSession, ID: "sess-heal", Reason: "x", RevokedAt: time.Now(), ExpiresAt: time.Now().Add(time.Hour)}
dead := Revocation{Kind: KindSession, ID: "sess-dead", Reason: "x", RevokedAt: time.Now().Add(-2 * time.Hour), ExpiresAt: time.Now().Add(-time.Hour)}
_ = fake.Upsert(ctx, live)
_ = fake.Upsert(ctx, dead)
s := New(Options{Durable: fake})
// Local cache starts empty (as after a Redis flush + local reconcile miss).
if s.IsRevoked("sess-heal", 0, time.Time{}) {
t.Fatalf("did not expect sess-heal in a fresh empty cache")
}
s.maintain(ctx)
if got := fake.pruneCount(); got != 1 {
t.Fatalf("durable prune calls = %d, want 1", got)
}
if !s.IsRevoked("sess-heal", 0, time.Time{}) {
t.Fatalf("expected sess-heal re-warmed into cache by maintain")
}
fake.mu.Lock()
_, deadStillThere := fake.rows[dead.key()]
fake.mu.Unlock()
if deadStillThere {
t.Fatalf("expected expired durable row to be pruned")
}
}
func TestIsRevoked(t *testing.T) {
ctx := context.Background()
// preCutoff stands in for a stream credential issued before any revocation
// a test writes; user-kind matching requires the credential to predate the
// revocation (the cutoff), session-kind matching ignores it.
preCutoff := time.Now().Add(-time.Minute)
tests := []struct {
name string
setup func(t *testing.T, s *Store)
sessionID string
userID int
startedAt time.Time
wantRevoke bool
}{
{
name: "revoked session is revoked",
setup: func(t *testing.T, s *Store) {
if err := s.RevokeSession(ctx, "sess-1", "abuse"); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
},
sessionID: "sess-1",
userID: 42,
startedAt: preCutoff,
wantRevoke: true,
},
{
name: "revoked user revokes that user's pre-revocation streams",
setup: func(t *testing.T, s *Store) {
if err := s.RevokeUser(ctx, 7, "banned"); err != nil {
t.Fatalf("RevokeUser: %v", err)
}
},
sessionID: "any-session-for-user-7",
userID: 7,
startedAt: preCutoff,
wantRevoke: true,
},
{
name: "user revocation is a cutoff: a post-revocation credential plays",
setup: func(t *testing.T, s *Store) {
if err := s.RevokeUser(ctx, 7, "sessions_revoked"); err != nil {
t.Fatalf("RevokeUser: %v", err)
}
},
sessionID: "fresh-session-for-user-7",
userID: 7,
startedAt: time.Now().Add(time.Second),
wantRevoke: false,
},
{
name: "user revocation never matches an unknown credential time (fail open)",
setup: func(t *testing.T, s *Store) {
if err := s.RevokeUser(ctx, 7, "banned"); err != nil {
t.Fatalf("RevokeUser: %v", err)
}
},
sessionID: "sess-without-start-time",
userID: 7,
startedAt: time.Time{},
wantRevoke: false,
},
{
name: "session revocation ignores the credential time",
setup: func(t *testing.T, s *Store) {
if err := s.RevokeSession(ctx, "sess-exact", "admin_terminate"); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
},
sessionID: "sess-exact",
userID: 42,
startedAt: time.Now().Add(time.Hour),
wantRevoke: true,
},
{
name: "expired revocation is not revoked",
setup: func(t *testing.T, s *Store) {
past := time.Now().Add(-time.Minute)
if err := s.Revoke(ctx, Key{Kind: KindSession, ID: "sess-old"}, "stale", past); err != nil {
t.Fatalf("Revoke: %v", err)
}
},
sessionID: "sess-old",
userID: 1,
startedAt: preCutoff,
wantRevoke: false,
},
{
name: "unrelated session and user are not revoked",
setup: func(t *testing.T, s *Store) {},
sessionID: "unknown",
userID: 999,
startedAt: preCutoff,
wantRevoke: false,
},
{
name: "unrelated session with a different revoked user is not revoked",
setup: func(t *testing.T, s *Store) {
if err := s.RevokeUser(ctx, 7, "banned"); err != nil {
t.Fatalf("RevokeUser: %v", err)
}
},
sessionID: "sess-other",
userID: 8,
startedAt: preCutoff,
wantRevoke: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
s := newMemStore()
tt.setup(t, s)
if got := s.IsRevoked(tt.sessionID, tt.userID, tt.startedAt); got != tt.wantRevoke {
t.Fatalf("IsRevoked(%q, %d, %v) = %v, want %v", tt.sessionID, tt.userID, tt.startedAt, got, tt.wantRevoke)
}
})
}
}
func TestList(t *testing.T) {
ctx := context.Background()
s := newMemStore()
if got := s.List(); len(got) != 0 {
t.Fatalf("List on empty store = %d entries, want 0", len(got))
}
if err := s.RevokeSession(ctx, "sess-a", "r1"); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
if err := s.RevokeUser(ctx, 5, "r2"); err != nil {
t.Fatalf("RevokeUser: %v", err)
}
// An expired entry must not appear in List.
if err := s.Revoke(ctx, Key{Kind: KindSession, ID: "sess-exp"}, "r3", time.Now().Add(-time.Second)); err != nil {
t.Fatalf("Revoke: %v", err)
}
got := s.List()
if len(got) != 2 {
t.Fatalf("List = %d active entries, want 2: %+v", len(got), got)
}
seen := make(map[Key]Revocation, len(got))
for _, r := range got {
seen[Key{Kind: r.Kind, ID: r.ID}] = r
}
if _, ok := seen[Key{Kind: KindSession, ID: "sess-a"}]; !ok {
t.Errorf("List missing revoked session sess-a")
}
if _, ok := seen[Key{Kind: KindUser, ID: "5"}]; !ok {
t.Errorf("List missing revoked user 5")
}
if _, ok := seen[Key{Kind: KindSession, ID: "sess-exp"}]; ok {
t.Errorf("List returned expired entry sess-exp")
}
}
func TestExpiryPrunesFromCache(t *testing.T) {
ctx := context.Background()
s := newMemStore()
// Revoke with a short future TTL, then confirm it lapses.
until := time.Now().Add(30 * time.Millisecond)
if err := s.Revoke(ctx, Key{Kind: KindUser, ID: "3"}, "temp", until); err != nil {
t.Fatalf("Revoke: %v", err)
}
preCutoff := time.Now().Add(-time.Minute)
if !s.IsRevoked("whatever", 3, preCutoff) {
t.Fatalf("expected user 3 to be revoked before expiry")
}
time.Sleep(50 * time.Millisecond)
if s.IsRevoked("whatever", 3, preCutoff) {
t.Fatalf("expected user 3 to no longer be revoked after expiry")
}
if got := s.List(); len(got) != 0 {
t.Fatalf("List after expiry = %d entries, want 0", len(got))
}
}
// TestMonotonicExpiry guards the invariant that a re-revoke with a SHORTER TTL
// can never shorten a longer-lived kill. The async over-cap enforcer re-revokes
// with a short self-healing TTL; without monotonic expiry it would shrink an
// admin's 24h RevokeSession on the same session key and reopen the
// restart-resurrection window.
func TestMonotonicExpiry(t *testing.T) {
ctx := context.Background()
s := newMemStore()
long := time.Now().Add(24 * time.Hour)
if err := s.Revoke(ctx, Key{Kind: KindSession, ID: "sess-1"}, "admin_terminate", long); err != nil {
t.Fatalf("Revoke long: %v", err)
}
// A shorter re-revoke (the enforcer's 5m self-heal TTL) must not win.
short := time.Now().Add(5 * time.Minute)
if err := s.Revoke(ctx, Key{Kind: KindSession, ID: "sess-1"}, "over_cap", short); err != nil {
t.Fatalf("Revoke short: %v", err)
}
got := s.List()
if len(got) != 1 {
t.Fatalf("List = %d entries, want 1", len(got))
}
if !got[0].ExpiresAt.Equal(long) {
t.Fatalf("ExpiresAt = %v, want the longer %v (shorter re-revoke must not shorten the kill)", got[0].ExpiresAt, long)
}
// A LONGER re-revoke does extend the kill.
longer := time.Now().Add(48 * time.Hour)
if err := s.Revoke(ctx, Key{Kind: KindSession, ID: "sess-1"}, "extended", longer); err != nil {
t.Fatalf("Revoke longer: %v", err)
}
got = s.List()
if len(got) != 1 || !got[0].ExpiresAt.Equal(longer) {
t.Fatalf("ExpiresAt after longer re-revoke = %v, want %v", got[0].ExpiresAt, longer)
}
}
// TestUserZeroNeverMatches guards the "no resolved owner" sentinel: a session-
// only check (userID 0, e.g. from the transcode node) must never be caught by a
// stray user:"0" revocation, which would read as "every ownerless request is
// revoked".
func TestUserZeroNeverMatches(t *testing.T) {
ctx := context.Background()
s := newMemStore()
// RevokeUser(0) is rejected outright: IsRevoked never matches a user:0 entry,
// so creating one would be a silently-ineffective kill. Guarding here means no
// stray user:0 revocation can exist to be misread as "every ownerless request
// is revoked".
if err := s.RevokeUser(ctx, 0, "should-not-nuke-everything"); err == nil {
t.Fatalf("RevokeUser(0) must be rejected, got nil error")
}
if s.IsRevoked("some-unrelated-session", 0, time.Now().Add(-time.Minute)) {
t.Fatalf("userID 0 sentinel must not match a user:0 revocation")
}
// A real session revocation still works with a 0 owner id.
if err := s.RevokeSession(ctx, "sess-x", "abuse"); err != nil {
t.Fatalf("RevokeSession: %v", err)
}
if !s.IsRevoked("sess-x", 0, time.Time{}) {
t.Fatalf("expected sess-x to be revoked even with owner id 0")
}
}
+11
View File
@@ -76,6 +76,17 @@ type Claims struct {
jwt.RegisteredClaims
}
// IssuedTime returns the token's iat as a time.Time, zero when absent. It is
// the credential-issue instant the revocation store's user-kill cutoff compares
// against (streamrevoke.Store.IsRevoked): a user revocation kills streams whose
// token predates it and spares ones minted after re-authentication.
func (c *Claims) IssuedTime() time.Time {
if c == nil || c.IssuedAt == nil {
return time.Time{}
}
return c.IssuedAt.Time
}
// Sign creates a signed JWT string from the given claims.
func Sign(c Claims, secret string, ttl time.Duration) (string, error) {
now := time.Now()
+39
View File
@@ -98,6 +98,7 @@ type Server struct {
reaperOnce sync.Once
mu sync.RWMutex
activeJobs atomic.Int32
revocation revocationStore
// reconstructGroup single-flights node-side session reconstruction per session
// id so a post-restart wave of concurrent manifest/segment requests for the same
@@ -378,6 +379,31 @@ func (s *Server) SetRecipeStore(store recipeStore) {
s.recipeStore = store
}
// revocationStore is the subset of *streamrevoke.Store the node consults to
// refuse serving (and rebuilding) killed sessions. IsRevoked is a pure in-memory
// lookup. Nil disables enforcement. Defense-in-depth: the proxy in front already
// enforces, but this guards the reconstruct path so a killed session is never
// re-spawned after a node restart.
type revocationStore interface {
IsRevoked(sessionID string, userID int, startedAt time.Time) bool
Refuse(w http.ResponseWriter, sessionID string, userID int, startedAt time.Time) bool
}
// SetRevocationStore wires the kill switch this node consults per request.
func (s *Server) SetRevocationStore(store revocationStore) {
s.revocation = store
}
// refuseIfRevoked refuses (403) a serve request for a revoked session. It checks
// the session key only (no per-segment token parse): the node is always fronted
// by the proxy/central, which enforce per-user revocations, and the enforcer +
// admin terminate revoke by session id. The reconstruct guard uses the verified
// token's user id directly, so per-user kills still block a rebuild after a
// node restart.
func (s *Server) refuseIfRevoked(w http.ResponseWriter, sessionID string) bool {
return s.revocation != nil && s.revocation.Refuse(w, sessionID, 0, time.Time{})
}
// Handler returns the chi.Router with all transcode routes.
func (s *Server) Handler() http.Handler {
s.startIdleReaper()
@@ -623,6 +649,11 @@ func (s *Server) reconstructFromToken(r *http.Request, sessionID string, request
"session", sessionID, "playback_session_id", sessionID)
return nil
}
// Never rebuild a killed session: without this guard a node restart would let
// reconstruction re-spawn ffmpeg for a session the kill switch already revoked.
if s.revocation != nil && s.revocation.IsRevoked(claims.SessionID, claims.UserID, claims.IssuedTime()) {
return nil
}
card := playback.RecipeCardFromClaims(claims)
// The token's recipe must be a transcode card for the session id in the URL: a
// mismatch is a forged or stale request, and direct/remux cards carry no encode
@@ -841,6 +872,10 @@ func (s *Server) handleStop(w http.ResponseWriter, r *http.Request) {
func (s *Server) handleManifest(w http.ResponseWriter, r *http.Request) {
sessionID := chi.URLParam(r, "session_id")
if s.refuseIfRevoked(w, sessionID) {
return
}
// Lookup and liveness refresh happen atomically so the idle reaper can
// never unregister the job between them and tear down a session this
// request is about to serve from.
@@ -880,6 +915,10 @@ func (s *Server) handleSegment(w http.ResponseWriter, r *http.Request) {
sessionID := chi.URLParam(r, "session_id")
name := chi.URLParam(r, "name")
if s.refuseIfRevoked(w, sessionID) {
return
}
// Lookup and liveness refresh happen atomically so the idle reaper can
// never unregister the job between them and tear down a session this
// request is about to serve from.
@@ -0,0 +1,20 @@
-- +goose Up
-- +goose StatementBegin
CREATE TABLE public.stream_revocations (
kind text NOT NULL,
id text NOT NULL,
reason text NOT NULL DEFAULT '',
revoked_at timestamptz NOT NULL DEFAULT now(),
expires_at timestamptz NOT NULL,
CONSTRAINT stream_revocations_pkey PRIMARY KEY (kind, id),
CONSTRAINT stream_revocations_kind_check CHECK (kind IN ('sess', 'user'))
);
CREATE INDEX stream_revocations_expires_at_idx
ON public.stream_revocations (expires_at);
-- +goose StatementEnd
-- +goose Down
-- +goose StatementBegin
DROP TABLE IF EXISTS public.stream_revocations;
-- +goose StatementEnd