diff --git a/internal/httpstream/rolling_deadline.go b/internal/httpstream/rolling_deadline.go index f406aa4a..c69a7472 100644 --- a/internal/httpstream/rolling_deadline.go +++ b/internal/httpstream/rolling_deadline.go @@ -17,6 +17,7 @@ import ( "net/http" "os" "strconv" + "sync/atomic" "time" ) @@ -70,30 +71,76 @@ type RollingDeadlineWriter struct { step time.Duration lastBump time.Time disabled bool + latch *CutLatch statusCode int bytesWritten int64 firstWriteErr error } +// CutLatch records that a stream has been terminally cut. Once latched, a +// RollingDeadlineWriter must never push the write deadline back out again: the +// cut is a deliberate hang-up, not a stall. +type CutLatch struct { + cut atomic.Bool +} + +func (l *CutLatch) Cut() { + if l != nil { + l.cut.Store(true) + } +} + +func (l *CutLatch) IsCut() bool { + return l != nil && l.cut.Load() +} + +type cutLatchContextKey struct{} + +// WithCutLatch carries l on ctx so rolling writers constructed inside serving +// helpers can observe a cut made by a watcher around an inner writer. +func WithCutLatch(ctx context.Context, l *CutLatch) context.Context { + return context.WithValue(ctx, cutLatchContextKey{}, l) +} + +// CutLatchFrom returns the stream cut latch carried by ctx, if any. +func CutLatchFrom(ctx context.Context) *CutLatch { + if ctx == nil { + return nil + } + l, _ := ctx.Value(cutLatchContextKey{}).(*CutLatch) + return l +} + // NewRollingDeadlineWriter wraps w with the configured stall window. func NewRollingDeadlineWriter(w http.ResponseWriter) *RollingDeadlineWriter { - return newRollingDeadlineWriter(w, StallWindow(), bumpStep) + return newRollingDeadlineWriterWithLatch(w, StallWindow(), bumpStep, nil) +} + +// NewRollingDeadlineWriterCtx wraps w and observes a terminal cut latch carried +// by ctx. Callers without a revocable request can use NewRollingDeadlineWriter. +func NewRollingDeadlineWriterCtx(ctx context.Context, w http.ResponseWriter) *RollingDeadlineWriter { + return newRollingDeadlineWriterWithLatch(w, StallWindow(), bumpStep, CutLatchFrom(ctx)) } func newRollingDeadlineWriter(w http.ResponseWriter, window, step time.Duration) *RollingDeadlineWriter { + return newRollingDeadlineWriterWithLatch(w, window, step, nil) +} + +func newRollingDeadlineWriterWithLatch(w http.ResponseWriter, window, step time.Duration, latch *CutLatch) *RollingDeadlineWriter { s := &RollingDeadlineWriter{ w: w, rc: http.NewResponseController(w), window: window, step: step, + latch: latch, } s.bump() return s } func (s *RollingDeadlineWriter) bump() { - if s.disabled { + if s.disabled || s.latch.IsCut() { return } now := time.Now() @@ -104,6 +151,13 @@ func (s *RollingDeadlineWriter) bump() { s.disabled = true return } + // Close the check/set race with a concurrent cut: if the watcher latched + // after the first check but before the future deadline landed, immediately + // restore the terminal deadline instead of leaving the socket re-armed. + if s.latch.IsCut() { + _ = s.rc.SetWriteDeadline(time.Now()) + return + } s.lastBump = now } diff --git a/internal/httpstream/rolling_deadline_test.go b/internal/httpstream/rolling_deadline_test.go index d64c548d..1300e2e6 100644 --- a/internal/httpstream/rolling_deadline_test.go +++ b/internal/httpstream/rolling_deadline_test.go @@ -211,6 +211,49 @@ func TestWriteOutcomeCompletedAndCounted(t *testing.T) { } } +type recordingDeadlineWriter struct { + *httptest.ResponseRecorder + deadlines []time.Time +} + +func (w *recordingDeadlineWriter) SetWriteDeadline(deadline time.Time) error { + w.deadlines = append(w.deadlines, deadline) + return nil +} + +func TestRollingDeadlineWriterNeverRearmsAfterCut(t *testing.T) { + base := &recordingDeadlineWriter{ResponseRecorder: httptest.NewRecorder()} + latch := &CutLatch{} + sw := newRollingDeadlineWriterWithLatch(base, time.Minute, 0, latch) + if got := len(base.deadlines); got != 1 { + t.Fatalf("constructor deadlines = %d, want 1", got) + } + + latch.Cut() + sw.bump() + if _, err := sw.Write([]byte("post-cut")); err != nil { + t.Fatal(err) + } + if got := len(base.deadlines); got != 1 { + t.Fatalf("deadlines after cut = %d, want constructor deadline only", got) + } +} + +func TestRollingDeadlineWriterStartsLatchedWithoutBump(t *testing.T) { + base := &recordingDeadlineWriter{ResponseRecorder: httptest.NewRecorder()} + latch := &CutLatch{} + latch.Cut() + ctx := WithCutLatch(context.Background(), latch) + + sw := newRollingDeadlineWriterWithLatch(base, time.Minute, 0, CutLatchFrom(ctx)) + if _, err := sw.Write([]byte("still no bump")); err != nil { + t.Fatal(err) + } + if got := len(base.deadlines); got != 0 { + t.Fatalf("deadlines = %d, want 0 for pre-latched writer", got) + } +} + func TestServeContentReadFromOutcomeCompletedAndCounted(t *testing.T) { const totalSize = 2 << 20 filePath := filepath.Join(t.TempDir(), "source.bin") diff --git a/internal/playback/directplay.go b/internal/playback/directplay.go index 8a211f2a..9b260a0f 100644 --- a/internal/playback/directplay.go +++ b/internal/playback/directplay.go @@ -57,7 +57,7 @@ func MimeFromExtension(name string) string { func ServeDirectPlay(w http.ResponseWriter, r *http.Request, filePath string) error { // Media bodies routinely take longer than the server's absolute // WriteTimeout; roll the write deadline with progress instead. - streamWriter := httpstream.NewRollingDeadlineWriter(w) + streamWriter := httpstream.NewRollingDeadlineWriterCtx(r.Context(), w) w = streamWriter f, err := os.Open(filePath) if err != nil { diff --git a/internal/playback/remux.go b/internal/playback/remux.go index 4f5d330e..ccda1de5 100644 --- a/internal/playback/remux.go +++ b/internal/playback/remux.go @@ -335,7 +335,7 @@ func ServeRemux(w http.ResponseWriter, r *http.Request, filePath, outputFormat s func ServeRemuxWithDVMode(w http.ResponseWriter, r *http.Request, filePath, outputFormat string, seekSeconds float64, transcodeAudio bool, audioTrackIndex int, dvProfile int, mode RemuxDVMode, ffmpegPath string) error { // Remux output streams for the length of the title; roll the write // deadline with progress instead of the server's absolute WriteTimeout. - w = httpstream.NewRollingDeadlineWriter(w) + w = httpstream.NewRollingDeadlineWriterCtx(r.Context(), w) // Check file exists before starting ffmpeg to return a proper 404. // Headers must be written before streaming begins, so we can't detect // ffmpeg errors after WriteHeader(200) has been sent. diff --git a/internal/streamrevoke/store.go b/internal/streamrevoke/store.go index 1864e24e..4863cdcb 100644 --- a/internal/streamrevoke/store.go +++ b/internal/streamrevoke/store.go @@ -21,6 +21,7 @@ import ( "time" "github.com/Silo-Server/silo-server/internal/cache" + "github.com/Silo-Server/silo-server/internal/httpstream" "github.com/redis/go-redis/v9" ) @@ -96,20 +97,22 @@ type DurableStore interface { // 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 + 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 + WatchInterval time.Duration // default 5s + 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 + rdb *redis.Client + bus cache.EventBus + durable DurableStore + pollInterval time.Duration + defaultTTL time.Duration + watchInterval time.Duration opMu sync.Mutex mu sync.RWMutex @@ -126,12 +129,16 @@ func New(opts Options) *Store { if opts.DefaultTTL <= 0 { opts.DefaultTTL = defaultTTL } + if opts.WatchInterval <= 0 { + opts.WatchInterval = 5 * time.Second + } return &Store{ rdb: opts.Redis, bus: opts.Bus, durable: opts.Durable, pollInterval: opts.PollInterval, defaultTTL: opts.DefaultTTL, + watchInterval: opts.WatchInterval, items: make(map[Key]Revocation), tombstones: make(map[Key]time.Time), tombstoneExpires: make(map[Key]time.Time), @@ -174,24 +181,44 @@ func (s *Store) Refuse(w http.ResponseWriter, sessionID string, userID int, star // 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. +// If the ResponseWriter chain doesn't support write deadlines, the failure is +// logged and the stream still stops on its next request via Refuse. A context +// cut latch also prevents rolling deadline writers from re-arming the socket. +// This helper 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() { + return s.WatchAndCutContext(context.Background(), w, sessionID, userID, startedAt) +} + +// WatchAndCutContext is WatchAndCut with the request context used for logging +// and for resolving the rolling-deadline cut latch. +func (s *Store) WatchAndCutContext(ctx context.Context, w http.ResponseWriter, sessionID string, userID int, startedAt time.Time) func() { if s == nil { return func() {} } - cut := func() { _ = http.NewResponseController(w).SetWriteDeadline(time.Now()) } + latch := httpstream.CutLatchFrom(ctx) + var warnOnce sync.Once + cut := func() { + latch.Cut() + if err := http.NewResponseController(w).SetWriteDeadline(time.Now()); err != nil { + warnOnce.Do(func() { + slog.WarnContext(ctx, "stream cut could not set write deadline; in-flight pour continues until its next request", + "component", "streamrevoke", + "session", sessionID, + "user", userID, + "error", err, + ) + }) + } + } if s.IsRevoked(sessionID, userID, startedAt) { cut() - return func() {} } done := make(chan struct{}) go func() { - ticker := time.NewTicker(5 * time.Second) + ticker := time.NewTicker(s.watchInterval) defer ticker.Stop() for { select { @@ -200,7 +227,6 @@ func (s *Store) WatchAndCut(w http.ResponseWriter, sessionID string, userID int, case <-ticker.C: if s.IsRevoked(sessionID, userID, startedAt) { cut() - return } } } diff --git a/internal/streamrevoke/store_test.go b/internal/streamrevoke/store_test.go index 664665d8..6c821b3c 100644 --- a/internal/streamrevoke/store_test.go +++ b/internal/streamrevoke/store_test.go @@ -1,14 +1,20 @@ package streamrevoke import ( + "bytes" "context" "encoding/json" "errors" + "log/slog" + "net/http" + "net/http/httptest" + "strings" "sync" "testing" "time" "github.com/Silo-Server/silo-server/internal/cache" + "github.com/Silo-Server/silo-server/internal/httpstream" "github.com/redis/go-redis/v9" ) @@ -20,6 +26,107 @@ func newMemStore() *Store { return New(Options{}) } +type cutDeadlineWriter struct { + mu sync.Mutex + header http.Header + deadlines []time.Time +} + +func (w *cutDeadlineWriter) Header() http.Header { + if w.header == nil { + w.header = make(http.Header) + } + return w.header +} + +func (w *cutDeadlineWriter) Write(p []byte) (int, error) { return len(p), nil } +func (w *cutDeadlineWriter) WriteHeader(int) {} +func (w *cutDeadlineWriter) SetWriteDeadline(deadline time.Time) error { + w.mu.Lock() + defer w.mu.Unlock() + w.deadlines = append(w.deadlines, deadline) + return nil +} + +func (w *cutDeadlineWriter) deadlineCount() int { + w.mu.Lock() + defer w.mu.Unlock() + return len(w.deadlines) +} + +func TestImmediateCutLatchesBeforeRollingWriterConstruction(t *testing.T) { + s := newMemStore() + if err := s.RevokeSession(context.Background(), "already-cut", "test"); err != nil { + t.Fatal(err) + } + latch := &httpstream.CutLatch{} + ctx := httpstream.WithCutLatch(context.Background(), latch) + base := &cutDeadlineWriter{} + + stop := s.WatchAndCutContext(ctx, base, "already-cut", 1, time.Now()) + defer stop() + if !latch.IsCut() { + t.Fatal("immediate revocation did not latch the terminal cut") + } + if got := base.deadlineCount(); got != 1 { + t.Fatalf("cut deadlines = %d, want 1", got) + } + + rolling := httpstream.NewRollingDeadlineWriterCtx(ctx, base) + if _, err := rolling.Write([]byte("must not rearm")); err != nil { + t.Fatal(err) + } + if got := base.deadlineCount(); got != 1 { + t.Fatalf("deadlines after rolling writer = %d, want cut only", got) + } +} + +func TestWatchAndCutKeepsReapplyingUntilStopped(t *testing.T) { + const watchInterval = 5 * time.Millisecond + s := New(Options{WatchInterval: watchInterval}) + if err := s.RevokeSession(context.Background(), "keep-cut", "test"); err != nil { + t.Fatal(err) + } + base := &cutDeadlineWriter{} + stop := s.WatchAndCutContext(context.Background(), base, "keep-cut", 1, time.Now()) + + deadline := time.Now().Add(time.Second) + for base.deadlineCount() < 3 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if got := base.deadlineCount(); got < 3 { + stop() + t.Fatalf("deadline applications = %d, want immediate cut plus ticker reapplications", got) + } + stop() + stoppedAt := base.deadlineCount() + time.Sleep(3 * watchInterval) + if got := base.deadlineCount(); got != stoppedAt { + t.Fatalf("deadline applications after stop = %d, want %d", got, stoppedAt) + } +} + +func TestWatchAndCutLogsUnsupportedDeadlineOnlyOnce(t *testing.T) { + const watchInterval = 5 * time.Millisecond + s := New(Options{WatchInterval: watchInterval}) + if err := s.RevokeSession(context.Background(), "unsupported-cut", "test"); err != nil { + t.Fatal(err) + } + var logs bytes.Buffer + previousLogger := slog.Default() + slog.SetDefault(slog.New(slog.NewTextHandler(&logs, nil))) + t.Cleanup(func() { slog.SetDefault(previousLogger) }) + + stop := s.WatchAndCutContext(context.Background(), httptest.NewRecorder(), "unsupported-cut", 1, time.Now()) + time.Sleep(3 * watchInterval) + stop() + + const message = "stream cut could not set write deadline; in-flight pour continues until its next request" + if got := strings.Count(logs.String(), message); got != 1 { + t.Fatalf("warning count = %d, want 1; logs: %s", got, logs.String()) + } +} + // 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