diff --git a/cmd/silo/main.go b/cmd/silo/main.go index bc631fe5..5acca885 100644 --- a/cmd/silo/main.go +++ b/cmd/silo/main.go @@ -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 diff --git a/internal/api/handlers/admin_playback_control.go b/internal/api/handlers/admin_playback_control.go index 6ffb5e6e..eeb1c6cf 100644 --- a/internal/api/handlers/admin_playback_control.go +++ b/internal/api/handlers/admin_playback_control.go @@ -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 diff --git a/internal/api/handlers/downloads.go b/internal/api/handlers/downloads.go index 9790d204..2e8fafc4 100644 --- a/internal/api/handlers/downloads.go +++ b/internal/api/handlers/downloads.go @@ -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) diff --git a/internal/api/middleware/metrics.go b/internal/api/middleware/metrics.go index b2aa0127..e91cb8fd 100644 --- a/internal/api/middleware/metrics.go +++ b/internal/api/middleware/metrics.go @@ -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 } diff --git a/internal/api/middleware/request_logger.go b/internal/api/middleware/request_logger.go index 7a8f3fc7..5caf4df1 100644 --- a/internal/api/middleware/request_logger.go +++ b/internal/api/middleware/request_logger.go @@ -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 } diff --git a/internal/api/router.go b/internal/api/router.go index 0cc99e29..a420f1bf 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -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) + } +} diff --git a/internal/jellycompat/handlers_playback.go b/internal/jellycompat/handlers_playback.go index fa5f9dfa..0eb8d992 100644 --- a/internal/jellycompat/handlers_playback.go +++ b/internal/jellycompat/handlers_playback.go @@ -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 { diff --git a/internal/jellycompat/image_proxy_tags.go b/internal/jellycompat/image_proxy_tags.go index b91eb109..5936df93 100644 --- a/internal/jellycompat/image_proxy_tags.go +++ b/internal/jellycompat/image_proxy_tags.go @@ -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;") diff --git a/internal/jellycompat/logging.go b/internal/jellycompat/logging.go index 28136679..33607a3a 100644 --- a/internal/jellycompat/logging.go +++ b/internal/jellycompat/logging.go @@ -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 } diff --git a/internal/jellycompat/router.go b/internal/jellycompat/router.go index 577bdbd7..4378665d 100644 --- a/internal/jellycompat/router.go +++ b/internal/jellycompat/router.go @@ -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. diff --git a/internal/jellycompat/server.go b/internal/jellycompat/server.go index 59b41754..945614d6 100644 --- a/internal/jellycompat/server.go +++ b/internal/jellycompat/server.go @@ -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 diff --git a/internal/jellycompat/streams.go b/internal/jellycompat/streams.go index d87c188d..aca0d69c 100644 --- a/internal/jellycompat/streams.go +++ b/internal/jellycompat/streams.go @@ -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 diff --git a/internal/proxy/egress.go b/internal/proxy/egress.go index e1a9c2fa..29e3ac0c 100644 --- a/internal/proxy/egress.go +++ b/internal/proxy/egress.go @@ -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 { diff --git a/internal/proxy/server.go b/internal/proxy/server.go index 9c5562dc..09a07b56 100644 --- a/internal/proxy/server.go +++ b/internal/proxy/server.go @@ -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) { diff --git a/internal/proxy/session_byte_writer_test.go b/internal/proxy/session_byte_writer_test.go new file mode 100644 index 00000000..56472f3f --- /dev/null +++ b/internal/proxy/session_byte_writer_test.go @@ -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) + } +} diff --git a/internal/streamenforcer/enforcer.go b/internal/streamenforcer/enforcer.go new file mode 100644 index 00000000..e7eb3211 --- /dev/null +++ b/internal/streamenforcer/enforcer.go @@ -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:] +} diff --git a/internal/streamenforcer/enforcer_test.go b/internal/streamenforcer/enforcer_test.go new file mode 100644 index 00000000..cf1f8789 --- /dev/null +++ b/internal/streamenforcer/enforcer_test.go @@ -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) + } +} diff --git a/internal/streamrevoke/durable_postgres.go b/internal/streamrevoke/durable_postgres.go new file mode 100644 index 00000000..99c3d73a --- /dev/null +++ b/internal/streamrevoke/durable_postgres.go @@ -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 +} diff --git a/internal/streamrevoke/store.go b/internal/streamrevoke/store.go new file mode 100644 index 00000000..b8f08867 --- /dev/null +++ b/internal/streamrevoke/store.go @@ -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 + } + } +} diff --git a/internal/streamrevoke/store_test.go b/internal/streamrevoke/store_test.go new file mode 100644 index 00000000..31ab4b43 --- /dev/null +++ b/internal/streamrevoke/store_test.go @@ -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") + } +} diff --git a/internal/streamtoken/token.go b/internal/streamtoken/token.go index 09debaa8..3b7ce83e 100644 --- a/internal/streamtoken/token.go +++ b/internal/streamtoken/token.go @@ -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() diff --git a/internal/transcodenode/server.go b/internal/transcodenode/server.go index d990713a..d2920817 100644 --- a/internal/transcodenode/server.go +++ b/internal/transcodenode/server.go @@ -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. diff --git a/migrations/sql/20260705025758_stream_revocations.sql b/migrations/sql/20260705025758_stream_revocations.sql new file mode 100644 index 00000000..f00709ae --- /dev/null +++ b/migrations/sql/20260705025758_stream_revocations.sql @@ -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