diff --git a/cmd/silo/main.go b/cmd/silo/main.go index 0c5e8cd5..25fcab64 100644 --- a/cmd/silo/main.go +++ b/cmd/silo/main.go @@ -1744,11 +1744,11 @@ func main() { // 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 + // + progress) is shared with the admin session list via streammonitor 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 + return streammonitor.LiveLocalSessions(sessionMgr, nodeIdentity), nil }) var monitorSource streammonitor.Source = localSource if apiRedisClient != nil { diff --git a/internal/api/handlers/nodes.go b/internal/api/handlers/nodes.go index 0a733637..bc78864f 100644 --- a/internal/api/handlers/nodes.go +++ b/internal/api/handlers/nodes.go @@ -23,43 +23,6 @@ import ( "github.com/redis/go-redis/v9" ) -// LiveLocalSessions maps the in-process playback sessions into monitoring -// records (the same SessionInfo shape edge nodes write to Redis), so integrated -// single-node streams — which never touch Redis — are still visible to the -// monitor and the admin active-streams view. Shared by the enforcer's FuncSource -// and the admin session list so there is exactly one Session→record mapping. -func LiveLocalSessions(sm *playback.SessionManager, nodeName string) []nodesessions.SessionInfo { - if sm == nil { - return nil - } - live := sm.AllSessions() - out := make([]nodesessions.SessionInfo, 0, len(live)) - for _, s := range live { - lastServedAt := s.LastServedAt - if lastServedAt.IsZero() { - lastServedAt = s.LastActivityAt - } - out = append(out, nodesessions.SessionInfo{ - SessionID: s.ID, - NodeName: nodeName, - AuthUserID: s.UserID, - ProfileID: s.ProfileID, - Type: string(s.PlayMethod), - Route: s.Origin(), - MediaFileID: s.MediaFileID, - ClientIP: s.ClientIP, - ClientName: s.ClientName, - Position: s.Position, - Resolution: s.TargetResolution, - HWAccel: s.TranscodeHWAccel, - StartedAt: s.StartedAt.UTC().Format(time.RFC3339), - LastServedAt: lastServedAt.UTC().Format(time.RFC3339), - BytesServed: s.BytesServed, - }) - } - return out -} - // NodeRepository defines the operations the NodeHandler needs on the node store. type NodeRepository interface { List(ctx context.Context) ([]*nodepool.Node, error) @@ -422,7 +385,7 @@ func (h *NodeHandler) HandleListSessions(w http.ResponseWriter, r *http.Request) // Integrated single-node streams live only in the in-process session manager. // Include them in the unfiltered listing (a node_id filter targets an edge). if h.sessionMgr != nil && nodeFilter == "" { - infos = append(infos, LiveLocalSessions(h.sessionMgr, h.localNodeName)...) + infos = append(infos, streammonitor.LiveLocalSessions(h.sessionMgr, h.localNodeName)...) } sessions := []json.RawMessage{} @@ -453,6 +416,9 @@ func (h *NodeHandler) HandleListSessions(w http.ResponseWriter, r *http.Request) // deployment can serve the field as an empty list forever. Collapsing the two // would advertise download monitoring that is not actually running. type nodeSessionsCapabilitiesResponse struct { + // LogicalSessionID reports that session records may carry the stable logical + // identity associated with a replaceable transport generation. + LogicalSessionID bool `json:"logical_session_id"` // Transfers reports that the node-sessions payload carries a transfers array. Transfers bool `json:"transfers"` // TransfersActive reports that a transfer registry is wired on this server, @@ -466,8 +432,9 @@ type nodeSessionsCapabilitiesResponse struct { // endpoint's availability rather than being advertised from an unrelated route. func (h *NodeHandler) HandleGetNodeSessionsCapabilities(w http.ResponseWriter, _ *http.Request) { writeJSON(w, http.StatusOK, nodeSessionsCapabilitiesResponse{ - Transfers: true, - TransfersActive: h.transfers != nil, + LogicalSessionID: true, + Transfers: true, + TransfersActive: h.transfers != nil, }) } diff --git a/internal/api/handlers/nodes_test.go b/internal/api/handlers/nodes_test.go index fae509f9..51887c50 100644 --- a/internal/api/handlers/nodes_test.go +++ b/internal/api/handlers/nodes_test.go @@ -118,8 +118,8 @@ func TestNodeSessionsCapabilitiesSeparatesSchemaFromRuntimeWiring(t *testing.T) if err := json.Unmarshal(rec.Body.Bytes(), &body); err != nil { t.Fatalf("decode response: %v", err) } - if !body.Transfers || !body.TransfersActive { - t.Fatalf("capabilities = %+v, want both transfers and transfers_active true", body) + if !body.LogicalSessionID || !body.Transfers || !body.TransfersActive { + t.Fatalf("capabilities = %+v, want schema flags and transfers_active true", body) } }) @@ -135,6 +135,9 @@ func TestNodeSessionsCapabilitiesSeparatesSchemaFromRuntimeWiring(t *testing.T) if !body.Transfers { t.Error("transfers = false; the response shape carries the key regardless of wiring") } + if !body.LogicalSessionID { + t.Error("logical_session_id = false; SessionInfo schema supports the field regardless of wiring") + } if body.TransfersActive { t.Error("transfers_active = true with no registry wired; that advertises monitoring that is not running") } diff --git a/internal/api/handlers/playback.go b/internal/api/handlers/playback.go index f9f154c6..05a8087a 100644 --- a/internal/api/handlers/playback.go +++ b/internal/api/handlers/playback.go @@ -2589,6 +2589,7 @@ func (h *PlaybackHandler) HandleChangeAudioTrack(w http.ResponseWriter, r *http. } nodeReq := transcodenode.TranscodeStartRequest{ SessionID: restartTransportID, + LogicalSessionID: updatedSession.ID, InputPath: file.FilePath, SourceVideoCodec: file.CodecVideo, SeekSeconds: restartSeekSeconds, @@ -3265,6 +3266,7 @@ func (h *PlaybackHandler) HandleStartTranscode(w http.ResponseWriter, r *http.Re } nodeReq := transcodenode.TranscodeStartRequest{ SessionID: replacementTransportID, + LogicalSessionID: session.ID, InputPath: file.FilePath, SourceVideoCodec: file.CodecVideo, VideoBitstreamFilter: videoBitstreamFilter, diff --git a/internal/api/handlers/playback_v3.go b/internal/api/handlers/playback_v3.go index 1c8ca19c..9cde8038 100644 --- a/internal/api/handlers/playback_v3.go +++ b/internal/api/handlers/playback_v3.go @@ -766,7 +766,33 @@ func (h *PlaybackHandler) prepareRemoteTransportV3(r *http.Request, session *pla videoCodec = "copy" } seekSeconds, startSegment := configureHLSTimelineV3(result.Plan, videoCodec, 2, float64(file.Duration)) - req := transcodenode.TranscodeStartRequest{SessionID: transportID, InputPath: file.FilePath, SourceVideoCodec: file.CodecVideo, VideoBitstreamFilter: videoBitstreamFilterForPlanV3(result.Plan), SeekSeconds: seekSeconds, StartSegmentNumber: startSegment, TargetResolution: result.TargetResolution, TargetCodecVideo: videoCodec, TargetCodecAudio: result.TargetAudioCodec, TargetAudioChannels: result.TargetAudioChannels, TargetBitrateKbps: result.TargetBitrateKbps, SegmentDuration: 2, HWAccel: h.playbackConfig().HWAccel, AudioTrackIndex: plannedAudioTrackIndexV3(result, session.AudioTrackIndex), SubtitleTrackIndex: result.SubtitleTransportTrackIndex, SubtitleBurnIn: result.SubtitleBurnIn, SubtitleCodec: result.SubtitleCodec, TotalDuration: float64(file.Duration), RequireReady: true} + req := transcodenode.TranscodeStartRequest{ + SessionID: transportID, + LogicalSessionID: session.ID, + InputPath: file.FilePath, + SourceVideoCodec: file.CodecVideo, + VideoBitstreamFilter: videoBitstreamFilterForPlanV3(result.Plan), + SeekSeconds: seekSeconds, + StartSegmentNumber: startSegment, + TargetResolution: result.TargetResolution, + TargetCodecVideo: videoCodec, + TargetCodecAudio: result.TargetAudioCodec, + TargetAudioChannels: result.TargetAudioChannels, + TargetBitrateKbps: result.TargetBitrateKbps, + SegmentDuration: 2, + HWAccel: h.playbackConfig().HWAccel, + AudioTrackIndex: plannedAudioTrackIndexV3(result, session.AudioTrackIndex), + SubtitleTrackIndex: result.SubtitleTransportTrackIndex, + SubtitleBurnIn: result.SubtitleBurnIn, + SubtitleCodec: result.SubtitleCodec, + TotalDuration: float64(file.Duration), + RequireReady: true, + AuthUserID: session.UserID, + ProfileID: session.ProfileID, + MediaFileID: file.ID, + Route: session.Origin(), + ClientName: session.ClientName, + } nodeResp, status, err := h.startRemotePlaybackTransport(r.Context(), node.URL, req) if err != nil { // A timeout can fire after the node actually started the job; the diff --git a/internal/api/handlers/playback_v3_test.go b/internal/api/handlers/playback_v3_test.go index 4cbe529d..8a7b45db 100644 --- a/internal/api/handlers/playback_v3_test.go +++ b/internal/api/handlers/playback_v3_test.go @@ -1156,6 +1156,14 @@ func TestPrepareTransportV3RequiresRemoteManifestReadiness(t *testing.T) { if !startRequest.RequireReady { t.Fatal("protocol-v3 remote start did not require manifest readiness") } + file := v3HandlerFixtureFile(t) + if startRequest.LogicalSessionID != "session-ready" || + startRequest.AuthUserID != 7 || + startRequest.ProfileID != "profile-1" || + startRequest.MediaFileID != file.ID || + startRequest.Route != playback.OriginNative { + t.Fatalf("protocol-v3 remote monitoring attribution = %+v", startRequest) + } } func TestHandleStartPlaybackUnknownProtocolUsesLegacyBranch(t *testing.T) { diff --git a/internal/jellycompat/handlers_playback.go b/internal/jellycompat/handlers_playback.go index 156bc1f1..4f983324 100644 --- a/internal/jellycompat/handlers_playback.go +++ b/internal/jellycompat/handlers_playback.go @@ -464,6 +464,7 @@ func (h *PlaybackHandler) startRemoteTranscode( reqBody := transcodenode.TranscodeStartRequest{ SessionID: upstreamSessionID, + LogicalSessionID: playSessionID, InputPath: file.FilePath, SeekSeconds: initialSeekSeconds, StartSegmentNumber: startSegmentNumber, diff --git a/internal/nodesessions/tracker.go b/internal/nodesessions/tracker.go index 8619ca31..398d31ce 100644 --- a/internal/nodesessions/tracker.go +++ b/internal/nodesessions/tracker.go @@ -27,18 +27,22 @@ const ( // SessionInfo represents an active streaming session stored in Redis. type SessionInfo struct { - SessionID string `json:"session_id"` - NodeURL string `json:"node_url"` - NodeName string `json:"node_name"` - UserID string `json:"user_id,omitempty"` - MediaItemID string `json:"media_item_id,omitempty"` - MediaTitle string `json:"media_title,omitempty"` - Type string `json:"type"` // "direct_play", "remux", "transcode" - CodecVideo string `json:"codec_video,omitempty"` - CodecAudio string `json:"codec_audio,omitempty"` - Resolution string `json:"resolution,omitempty"` - HWAccel string `json:"hw_accel,omitempty"` - StartedAt string `json:"started_at"` + SessionID string `json:"session_id"` + // LogicalSessionID is the stable playback-session identity when SessionID is + // a replaceable transport generation. It is opaque monitoring metadata; + // SessionID remains the node-addressing key. + LogicalSessionID string `json:"logical_session_id,omitempty"` + NodeURL string `json:"node_url"` + NodeName string `json:"node_name"` + UserID string `json:"user_id,omitempty"` + MediaItemID string `json:"media_item_id,omitempty"` + MediaTitle string `json:"media_title,omitempty"` + Type string `json:"type"` // "direct_play", "remux", "transcode" + CodecVideo string `json:"codec_video,omitempty"` + CodecAudio string `json:"codec_audio,omitempty"` + Resolution string `json:"resolution,omitempty"` + HWAccel string `json:"hw_accel,omitempty"` + StartedAt string `json:"started_at"` // LastServedAt / BytesServed are the authoritative, server-observed liveness // signals: they are refreshed only when the node actually serves bytes for this @@ -69,6 +73,20 @@ type SessionInfo struct { Position float64 `json:"position,omitempty"` } +// Lease identifies one request-scoped reference to a tracked session. Its +// fields are deliberately private: only the Tracker that issued it can release +// it, and stale leases from an invalidated generation are harmless. +type Lease struct { + sessionID string + epoch uint64 + id uint64 +} + +type sessionRefs struct { + epoch uint64 + leases map[uint64]struct{} +} + // Tracker manages session lifecycle in Redis for a single node. type Tracker struct { rdb *redis.Client @@ -77,11 +95,14 @@ type Tracker struct { nodeType string nodeHash string // first 8 chars of SHA-256 of nodeURL - mu sync.Mutex - sessions map[string]struct{} // set of active (explicitly tracked) session IDs - touched map[string]time.Time // ephemeral sessions by last-activity time - records map[string]SessionInfo // last-written record per session, for enriched refresh - bytes map[string]int64 // cumulative bytes served per session (monitoring only) + mu sync.Mutex + sessions map[string]sessionRefs // explicitly tracked sessions and request leases + ephemeral map[string]time.Time // ephemeral sessions by last-observed request time + touched map[string]time.Time // last successful serve time + records map[string]SessionInfo // last-written record per session, for enriched refresh + bytes map[string]int64 // cumulative bytes served per session (monitoring only) + nextEpoch uint64 + nextLease uint64 } // NewTracker creates a session tracker for the given node. @@ -89,15 +110,16 @@ type Tracker struct { func NewTracker(rdb *redis.Client, nodeURL, nodeName, nodeType string) *Tracker { h := sha256.Sum256([]byte(nodeURL)) return &Tracker{ - rdb: rdb, - nodeURL: nodeURL, - nodeName: nodeName, - nodeType: nodeType, - nodeHash: hex.EncodeToString(h[:4]), // 8 hex chars - sessions: make(map[string]struct{}), - touched: make(map[string]time.Time), - records: make(map[string]SessionInfo), - bytes: make(map[string]int64), + rdb: rdb, + nodeURL: nodeURL, + nodeName: nodeName, + nodeType: nodeType, + nodeHash: hex.EncodeToString(h[:4]), // 8 hex chars + sessions: make(map[string]sessionRefs), + ephemeral: make(map[string]time.Time), + touched: make(map[string]time.Time), + records: make(map[string]SessionInfo), + bytes: make(map[string]int64), } } @@ -128,7 +150,7 @@ func (tr *Tracker) ActiveCount() int { defer tr.mu.Unlock() now := time.Now() count := len(tr.sessions) - for id, last := range tr.touched { + for id, last := range tr.ephemeral { if _, dup := tr.sessions[id]; dup { continue } @@ -151,8 +173,8 @@ func (tr *Tracker) Snapshot() []SessionInfo { for id, rec := range tr.records { // Session-backed entries are always live until Remove; only ephemeral // (non-session) entries age out by idle timeout. - if _, isSession := tr.sessions[id]; !isSession { - if last, ok := tr.touched[id]; ok && now.Sub(last) > idleTTLFor(rec) { + if refs, isSession := tr.sessions[id]; !isSession || len(refs.leases) == 0 { + if last, ok := tr.ephemeral[id]; !ok || now.Sub(last) > idleTTLFor(rec) { continue } } @@ -173,10 +195,12 @@ func (tr *Tracker) enrichLocked(id string, rec SessionInfo) SessionInfo { return rec } -// Track registers an active session in Redis with a TTL. -func (tr *Tracker) Track(ctx context.Context, info SessionInfo) { +// Track registers an active session in Redis with a TTL and returns a lease. +// Request-scoped callers must release it exactly once. Lifecycle-owned callers +// may instead use Remove for unconditional teardown. +func (tr *Tracker) Track(ctx context.Context, info SessionInfo) Lease { if tr.rdb == nil { - return + return Lease{} } now := time.Now() if info.LastServedAt == "" { @@ -190,22 +214,63 @@ func (tr *Tracker) Track(ctx context.Context, info SessionInfo) { // is open. tr.mu.Lock() tr.preserveStartedAtLocked(&info) - tr.sessions[info.SessionID] = struct{}{} + refs, exists := tr.sessions[info.SessionID] + if !exists { + tr.nextEpoch++ + refs = sessionRefs{epoch: tr.nextEpoch, leases: make(map[uint64]struct{})} + } + tr.nextLease++ + lease := Lease{sessionID: info.SessionID, epoch: refs.epoch, id: tr.nextLease} + refs.leases[lease.id] = struct{}{} + tr.sessions[info.SessionID] = refs + delete(tr.ephemeral, info.SessionID) tr.records[info.SessionID] = info tr.touched[info.SessionID] = now enriched := tr.enrichLocked(info.SessionID, info) tr.mu.Unlock() + if exists { + return lease + } data, err := json.Marshal(enriched) if err != nil { slog.DebugContext(ctx, "session track marshal failed", "component", "nodesessions", "error", err) - return + return lease } key := tr.redisKey(info.SessionID) if err := tr.rdb.Set(ctx, key, data, sessionTTL).Err(); err != nil { slog.DebugContext(ctx, "session track set failed", "component", "nodesessions", "error", err, "session", info.SessionID) + } + return lease +} + +// EnsureEphemeral makes a short-lived session visible without claiming that any +// bytes were served. The first write is synchronous; later liveness/byte updates +// are flushed by the refresh loop. +func (tr *Tracker) EnsureEphemeral(ctx context.Context, info SessionInfo) { + if tr.rdb == nil { return } + now := time.Now() + tr.mu.Lock() + tr.preserveStartedAtLocked(&info) + _, known := tr.records[info.SessionID] + tr.ephemeral[info.SessionID] = now + tr.records[info.SessionID] = info + enriched := tr.enrichLocked(info.SessionID, info) + tr.mu.Unlock() + if known { + return + } + + data, err := json.Marshal(enriched) + if err != nil { + slog.DebugContext(ctx, "session ensure marshal failed", "component", "nodesessions", "error", err) + return + } + if err := tr.rdb.Set(ctx, tr.redisKey(info.SessionID), data, sessionTTL).Err(); err != nil { + slog.DebugContext(ctx, "session ensure set failed", "component", "nodesessions", "error", err, "session", info.SessionID) + } } // Touch registers or refreshes an ephemeral session that has no explicit end, @@ -218,26 +283,12 @@ func (tr *Tracker) Touch(ctx context.Context, info SessionInfo) { if tr.rdb == nil { return } - now := time.Now() + tr.EnsureEphemeral(ctx, info) tr.mu.Lock() - tr.preserveStartedAtLocked(&info) - _, known := tr.touched[info.SessionID] - tr.touched[info.SessionID] = now - tr.records[info.SessionID] = info - enriched := tr.enrichLocked(info.SessionID, info) + if _, known := tr.records[info.SessionID]; known { + tr.touched[info.SessionID] = time.Now() + } tr.mu.Unlock() - if known { - return - } - - data, err := json.Marshal(enriched) - if err != nil { - slog.DebugContext(ctx, "session touch marshal failed", "component", "nodesessions", "error", err) - return - } - if err := tr.rdb.Set(ctx, tr.redisKey(info.SessionID), data, sessionTTL).Err(); err != nil { - slog.DebugContext(ctx, "session touch set failed", "component", "nodesessions", "error", err, "session", info.SessionID) - } } // preserveStartedAtLocked keeps the first-seen StartedAt when a record is @@ -291,13 +342,52 @@ func (tr *Tracker) MarkServed(sessionID string) { tr.mu.Unlock() } -// Remove deletes a session from Redis and the in-memory set. +// Release drops one request-scoped lease. A stale, duplicate, or already +// invalidated release is a no-op. The final live lease tears the record down. +func (tr *Tracker) Release(ctx context.Context, lease Lease) { + if tr.rdb == nil || lease.sessionID == "" { + return + } + tr.mu.Lock() + refs, ok := tr.sessions[lease.sessionID] + if !ok || refs.epoch != lease.epoch { + tr.mu.Unlock() + slog.DebugContext(ctx, "stale session lease release ignored", "component", "nodesessions", "session", lease.sessionID) + return + } + if _, ok := refs.leases[lease.id]; !ok { + tr.mu.Unlock() + slog.DebugContext(ctx, "duplicate session lease release ignored", "component", "nodesessions", "session", lease.sessionID) + return + } + delete(refs.leases, lease.id) + if len(refs.leases) > 0 { + tr.sessions[lease.sessionID] = refs + tr.mu.Unlock() + return + } + delete(tr.sessions, lease.sessionID) + delete(tr.ephemeral, lease.sessionID) + delete(tr.touched, lease.sessionID) + delete(tr.records, lease.sessionID) + delete(tr.bytes, lease.sessionID) + tr.mu.Unlock() + + if err := tr.rdb.Del(ctx, tr.redisKey(lease.sessionID)).Err(); err != nil { + slog.DebugContext(ctx, "session release failed", "component", "nodesessions", "error", err, "session", lease.sessionID) + } +} + +// Remove deletes a session unconditionally and invalidates every outstanding +// lease for its generation. func (tr *Tracker) Remove(ctx context.Context, sessionID string) { if tr.rdb == nil { return } tr.mu.Lock() + tr.nextEpoch++ delete(tr.sessions, sessionID) + delete(tr.ephemeral, sessionID) delete(tr.touched, sessionID) delete(tr.records, sessionID) delete(tr.bytes, sessionID) @@ -314,16 +404,18 @@ func (tr *Tracker) Cleanup(ctx context.Context) { return } tr.mu.Lock() - ids := make([]string, 0, len(tr.sessions)+len(tr.touched)) + ids := make([]string, 0, len(tr.sessions)+len(tr.ephemeral)) for id := range tr.sessions { ids = append(ids, id) } - for id := range tr.touched { + for id := range tr.ephemeral { if _, dup := tr.sessions[id]; !dup { ids = append(ids, id) } } - tr.sessions = make(map[string]struct{}) + tr.nextEpoch++ + tr.sessions = make(map[string]sessionRefs) + tr.ephemeral = make(map[string]time.Time) tr.touched = make(map[string]time.Time) tr.records = make(map[string]SessionInfo) tr.bytes = make(map[string]int64) @@ -373,12 +465,13 @@ func (tr *Tracker) refreshAll(ctx context.Context) { for id := range tr.sessions { ids = append(ids, id) } - for id, last := range tr.touched { + for id, last := range tr.ephemeral { _, isSession := tr.sessions[id] if !isSession && now.Sub(last) > idleTTLFor(tr.records[id]) { // Idle ephemeral session: stop refreshing and let the Redis key // expire on its own. Session-backed entries (direct/remux) are pruned // on Remove, never by idle timeout, so a quiet-but-open pour stays live. + delete(tr.ephemeral, id) delete(tr.touched, id) delete(tr.records, id) delete(tr.bytes, id) diff --git a/internal/nodesessions/tracker_idle_test.go b/internal/nodesessions/tracker_idle_test.go index 15c0695a..2034e953 100644 --- a/internal/nodesessions/tracker_idle_test.go +++ b/internal/nodesessions/tracker_idle_test.go @@ -13,9 +13,9 @@ func TestTranscodeIdleTTLAgreesAcrossCountSnapshotAndRefresh(t *testing.T) { t.Cleanup(func() { _ = rdb.Close() }) tr := NewTracker(rdb, "http://node", "node", "proxy") idle := time.Now().Add(-90 * time.Second) - tr.touched[sessionTypeTranscode] = idle + tr.ephemeral[sessionTypeTranscode] = idle tr.records[sessionTypeTranscode] = SessionInfo{SessionID: sessionTypeTranscode, Type: sessionTypeTranscode} - tr.touched["direct"] = idle + tr.ephemeral["direct"] = idle tr.records["direct"] = SessionInfo{SessionID: "direct", Type: "direct_play"} if got := tr.ActiveCount(); got != 1 { diff --git a/internal/nodesessions/tracker_lease_test.go b/internal/nodesessions/tracker_lease_test.go new file mode 100644 index 00000000..84a79e41 --- /dev/null +++ b/internal/nodesessions/tracker_lease_test.go @@ -0,0 +1,176 @@ +package nodesessions + +import ( + "context" + "encoding/json" + "net" + "sync" + "testing" + + "github.com/redis/go-redis/v9" +) + +type memoryRedisHook struct { + mu sync.Mutex + values map[string]string +} + +func (h *memoryRedisHook) DialHook(next redis.DialHook) redis.DialHook { return next } + +func (h *memoryRedisHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { + return func(ctx context.Context, cmd redis.Cmder) error { + h.mu.Lock() + defer h.mu.Unlock() + switch c := cmd.(type) { + case *redis.StatusCmd: + if c.Name() == "set" { + args := c.Args() + var value string + switch v := args[2].(type) { + case string: + value = v + case []byte: + value = string(v) + } + h.values[args[1].(string)] = value + c.SetVal("OK") + return nil + } + case *redis.IntCmd: + if c.Name() == "del" { + for _, arg := range c.Args()[1:] { + delete(h.values, arg.(string)) + } + c.SetVal(1) + return nil + } + } + return next(ctx, cmd) + } +} + +func (h *memoryRedisHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return func(ctx context.Context, cmds []redis.Cmder) error { + for _, cmd := range cmds { + if err := h.ProcessHook(func(context.Context, redis.Cmder) error { return nil })(ctx, cmd); err != nil { + return err + } + } + return nil + } +} + +func newMemoryTracker(t *testing.T) (*Tracker, *memoryRedisHook) { + t.Helper() + rdb := redis.NewClient(&redis.Options{ + Dialer: func(context.Context, string, string) (net.Conn, error) { + t.Fatal("unexpected Redis dial") + return nil, nil + }, + }) + t.Cleanup(func() { _ = rdb.Close() }) + hook := &memoryRedisHook{values: make(map[string]string)} + rdb.AddHook(hook) + return NewTracker(rdb, "http://node", "node", "proxy"), hook +} + +func (h *memoryRedisHook) has(key string) bool { + h.mu.Lock() + defer h.mu.Unlock() + _, ok := h.values[key] + return ok +} + +func TestStaleLeaseCannotReleaseReplacementAfterRemove(t *testing.T) { + tr, _ := newMemoryTracker(t) + ctx := context.Background() + old := tr.Track(ctx, SessionInfo{SessionID: "s"}) + tr.Remove(ctx, "s") + current := tr.Track(ctx, SessionInfo{SessionID: "s"}) + + tr.Release(ctx, old) + if got := tr.Snapshot(); len(got) != 1 || got[0].SessionID != "s" { + t.Fatalf("stale release removed replacement: %+v", got) + } + tr.Release(ctx, current) +} + +func TestStaleLeaseCannotReleaseReplacementAfterCleanup(t *testing.T) { + tr, _ := newMemoryTracker(t) + ctx := context.Background() + old := tr.Track(ctx, SessionInfo{SessionID: "s"}) + tr.Cleanup(ctx) + current := tr.Track(ctx, SessionInfo{SessionID: "s"}) + + tr.Release(ctx, old) + if got := tr.Snapshot(); len(got) != 1 || got[0].SessionID != "s" { + t.Fatalf("stale release after cleanup removed replacement: %+v", got) + } + tr.Release(ctx, current) +} + +func TestDoubleReleaseDoesNotConsumeAnotherLease(t *testing.T) { + tr, _ := newMemoryTracker(t) + ctx := context.Background() + first := tr.Track(ctx, SessionInfo{SessionID: "s"}) + second := tr.Track(ctx, SessionInfo{SessionID: "s"}) + + tr.Release(ctx, first) + tr.Release(ctx, first) + if got := tr.Snapshot(); len(got) != 1 { + t.Fatalf("double release consumed live lease: %+v", got) + } + tr.Release(ctx, second) +} + +func TestOverlappingLeasesKeepRecordKeyAndBytesAlive(t *testing.T) { + tr, redisState := newMemoryTracker(t) + ctx := context.Background() + first := tr.Track(ctx, SessionInfo{SessionID: "s"}) + second := tr.Track(ctx, SessionInfo{SessionID: "s"}) + + tr.Release(ctx, first) + tr.AddBytes("s", 123) + got := tr.Snapshot() + if len(got) != 1 || got[0].BytesServed != 123 { + t.Fatalf("live overlap record = %+v, want 123 bytes", got) + } + if !redisState.has(tr.redisKey("s")) { + t.Fatal("first release deleted Redis key while second lease was live") + } + + tr.Release(ctx, second) + if redisState.has(tr.redisKey("s")) { + t.Fatal("final release left Redis key behind") + } +} + +func TestEnsureEphemeralSeparatesVisibilityFromServedLiveness(t *testing.T) { + tr, redisState := newMemoryTracker(t) + ctx := context.Background() + tr.EnsureEphemeral(ctx, SessionInfo{SessionID: "s", Type: "transcode"}) + got := tr.Snapshot() + if len(got) != 1 || got[0].LastServedAt != "" || got[0].BytesServed != 0 { + t.Fatalf("ensured record claims served liveness: %+v", got) + } + if !redisState.has(tr.redisKey("s")) { + t.Fatal("ensured record was not projected to Redis") + } + + tr.AddBytes("s", 7) + got = tr.Snapshot() + if len(got) != 1 || got[0].LastServedAt == "" || got[0].BytesServed != 7 { + t.Fatalf("served record = %+v, want liveness and bytes", got) + } + + redisState.mu.Lock() + var projected SessionInfo + err := json.Unmarshal([]byte(redisState.values[tr.redisKey("s")]), &projected) + redisState.mu.Unlock() + if err != nil { + t.Fatalf("decode projected record: %v", err) + } + if projected.LastServedAt != "" { + t.Fatalf("first projection LastServedAt = %q, want empty", projected.LastServedAt) + } +} diff --git a/internal/proxy/remove_tracked_test.go b/internal/proxy/remove_tracked_test.go index ff03366d..2a5da30b 100644 --- a/internal/proxy/remove_tracked_test.go +++ b/internal/proxy/remove_tracked_test.go @@ -3,6 +3,7 @@ package proxy import ( "context" "net" + "sync" "testing" "github.com/Silo-Server/silo-server/internal/nodesessions" @@ -10,8 +11,10 @@ import ( ) type deleteCaptureHook struct { + mu sync.Mutex called bool contextErr error + values map[string]string } func (h *deleteCaptureHook) DialHook(next redis.DialHook) redis.DialHook { @@ -20,9 +23,25 @@ func (h *deleteCaptureHook) DialHook(next redis.DialHook) redis.DialHook { func (h *deleteCaptureHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { return func(ctx context.Context, cmd redis.Cmder) error { - if cmd.Name() == "del" { + h.mu.Lock() + defer h.mu.Unlock() + switch cmd.Name() { + case "set": + args := cmd.Args() + h.values[args[1].(string)] = "set" + if status, ok := cmd.(*redis.StatusCmd); ok { + status.SetVal("OK") + } + return nil + case "del": h.called = true h.contextErr = ctx.Err() + for _, arg := range cmd.Args()[1:] { + delete(h.values, arg.(string)) + } + if count, ok := cmd.(*redis.IntCmd); ok { + count.SetVal(1) + } return nil } return next(ctx, cmd) @@ -33,17 +52,31 @@ func (h *deleteCaptureHook) ProcessPipelineHook(next redis.ProcessPipelineHook) return next } -func TestRemoveTrackedUsesLiveContextAfterRequestCancellation(t *testing.T) { +func newProxyTestTracker(t *testing.T) (*nodesessions.Tracker, *deleteCaptureHook) { + t.Helper() rdb := redis.NewClient(&redis.Options{ Dialer: func(context.Context, string, string) (net.Conn, error) { - t.Fatal("DEL should be intercepted before dialing") + t.Fatal("Redis command should be intercepted before dialing") return nil, nil }, }) t.Cleanup(func() { _ = rdb.Close() }) - hook := &deleteCaptureHook{} + hook := &deleteCaptureHook{values: make(map[string]string)} rdb.AddHook(hook) - server := &Server{tracker: nodesessions.NewTracker(rdb, "http://node", "node", "proxy")} + return nodesessions.NewTracker(rdb, "http://node", "node", "proxy"), hook +} + +func (h *deleteCaptureHook) has(key string) bool { + h.mu.Lock() + defer h.mu.Unlock() + _, ok := h.values[key] + return ok +} + +func TestReleaseTrackedUsesLiveContextAfterRequestCancellation(t *testing.T) { + tracker, hook := newProxyTestTracker(t) + server := &Server{tracker: tracker} + lease := tracker.Track(context.Background(), nodesessions.SessionInfo{SessionID: "session-1"}) requestCtx, cancel := context.WithCancel(context.Background()) cancel() @@ -51,9 +84,11 @@ func TestRemoveTrackedUsesLiveContextAfterRequestCancellation(t *testing.T) { t.Fatal("request context was not canceled") } func() { - defer server.removeTracked("session-1") + defer server.releaseTracked(lease) }() + hook.mu.Lock() + defer hook.mu.Unlock() if !hook.called { t.Fatal("Redis DEL was not attempted") } diff --git a/internal/proxy/server.go b/internal/proxy/server.go index 709e53d6..a68f9cf5 100644 --- a/internal/proxy/server.go +++ b/internal/proxy/server.go @@ -181,14 +181,14 @@ func (s *Server) cutOnRevocation(ctx context.Context, w http.ResponseWriter, cla return s.revocation.WatchAndCutContext(ctx, w, claims.SessionID, claims.UserID, claims.IssuedTime()) } -// removeTracked deletes the edge record with a bounded background context. The +// releaseTracked releases one request lease with a bounded background context. The // request context is already canceled by the time this defer runs on a client // disconnect, which would skip the Redis DEL and leave a phantom session until // TTL — a false over-cap window. -func (s *Server) removeTracked(sessionID string) { +func (s *Server) releaseTracked(lease nodesessions.Lease) { ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() - s.tracker.Remove(ctx, sessionID) + s.tracker.Release(ctx, lease) } func (s *Server) handleDirectPlay(w http.ResponseWriter, r *http.Request) { @@ -199,8 +199,8 @@ func (s *Server) handleDirectPlay(w http.ResponseWriter, r *http.Request) { } info := sessionInfo(s.tracker, claims, "direct_play", edgeClientIP(r)) - s.tracker.Track(r.Context(), info) - defer s.removeTracked(claims.SessionID) + lease := s.tracker.Track(r.Context(), info) + defer s.releaseTracked(lease) sw := &sessionByteWriter{ResponseWriter: w, tracker: s.tracker, sessionID: claims.SessionID} defer sw.flush() @@ -292,8 +292,8 @@ func (s *Server) handleRemux(w http.ResponseWriter, r *http.Request) { } info := sessionInfo(s.tracker, claims, "remux", edgeClientIP(r)) - s.tracker.Track(r.Context(), info) - defer s.removeTracked(claims.SessionID) + lease := s.tracker.Track(r.Context(), info) + defer s.releaseTracked(lease) sw := &sessionByteWriter{ResponseWriter: w, tracker: s.tracker, sessionID: claims.SessionID} defer sw.flush() @@ -319,7 +319,7 @@ func (s *Server) handleTranscodeManifest(w http.ResponseWriter, r *http.Request) if claims == nil { return } - s.touchTranscodeSession(r, claims) + s.ensureTranscodeSession(r, claims) s.proxyToTranscodeNode(w, r, claims, "/transcode/"+transcodeTransportIDFromClaims(claims)+"/master.m3u8") } @@ -328,7 +328,7 @@ func (s *Server) handleTranscodeSegment(w http.ResponseWriter, r *http.Request) if claims == nil { return } - s.touchTranscodeSession(r, claims) + s.ensureTranscodeSession(r, claims) name := chi.URLParam(r, "name") s.proxyToTranscodeNode(w, r, claims, "/transcode/"+transcodeTransportIDFromClaims(claims)+"/segment/"+name) } @@ -343,12 +343,13 @@ func transcodeTransportIDFromClaims(claims *streamtoken.Claims) string { return claims.SessionID } -// touchTranscodeSession keeps HLS sessions visible in the active stream count. +// ensureTranscodeSession keeps HLS sessions visible in the active stream count. // Unlike direct play and remux, transcode playback reaches the proxy as many // 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", edgeClientIP(r))) +// requests instead of request lifetime. Visibility is separate from served +// liveness: only a successful upstream media response advances LastServedAt. +func (s *Server) ensureTranscodeSession(r *http.Request, claims *streamtoken.Claims) { + s.tracker.EnsureEphemeral(r.Context(), sessionInfo(s.tracker, claims, "transcode", edgeClientIP(r))) } // sessionInfo builds the node-session tracker record for a verified token, @@ -536,6 +537,10 @@ func (s *Server) proxyToTranscodeNode(w http.ResponseWriter, r *http.Request, cl } } w.WriteHeader(resp.StatusCode) + if resp.StatusCode != http.StatusOK && resp.StatusCode != http.StatusPartialContent { + _, _ = io.Copy(w, resp.Body) + return + } // 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 diff --git a/internal/proxy/tracker_lifecycle_test.go b/internal/proxy/tracker_lifecycle_test.go new file mode 100644 index 00000000..065bba6b --- /dev/null +++ b/internal/proxy/tracker_lifecycle_test.go @@ -0,0 +1,222 @@ +package proxy + +import ( + "bytes" + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "sync" + "testing" + "time" + + "github.com/Silo-Server/silo-server/internal/config" + "github.com/Silo-Server/silo-server/internal/nodeconfig" + "github.com/Silo-Server/silo-server/internal/nodesessions" + "github.com/Silo-Server/silo-server/internal/streamtoken" +) + +const proxyTrackerTestSecret = "proxy-tracker-test-secret" + +func newMountedProxyTestServer(t *testing.T, mediaPath string) (*Server, *deleteCaptureHook, string) { + t.Helper() + watcher := nodeconfig.NewWatcher(nil, nil, nil, nodeconfig.BootstrapOverrides{}) + cfg := &config.Config{} + cfg.Auth.JWTSecret = proxyTrackerTestSecret + watcher.SetConfigForTest(cfg) + tracker, redisState := newProxyTestTracker(t) + server := NewServer(watcher, tracker) + token, err := streamtoken.Sign(streamtoken.Claims{ + SessionID: "session-1", + MediaPath: mediaPath, + TranscodeNode: "http://transcode", + UserID: 7, + ProfileID: "profile-1", + MediaFileID: 42, + }, proxyTrackerTestSecret, time.Minute) + if err != nil { + t.Fatalf("sign stream token: %v", err) + } + return server, redisState, token +} + +type blockingResponseWriter struct { + header http.Header + entered chan struct{} + release chan struct{} + once sync.Once +} + +func (w *blockingResponseWriter) Header() http.Header { return w.header } +func (w *blockingResponseWriter) WriteHeader(int) {} +func (w *blockingResponseWriter) Write(p []byte) (int, error) { + w.once.Do(func() { close(w.entered) }) + <-w.release + return len(p), nil +} + +func TestMountedDirectOverlappingRangesKeepSessionUntilFinalRequest(t *testing.T) { + path := t.TempDir() + "/movie.bin" + if err := os.WriteFile(path, bytes.Repeat([]byte("x"), 64<<10), 0o600); err != nil { + t.Fatal(err) + } + server, redisState, token := newMountedProxyTestServer(t, path) + firstWriter := &blockingResponseWriter{ + header: make(http.Header), + entered: make(chan struct{}), + release: make(chan struct{}), + } + firstDone := make(chan struct{}) + go func() { + defer close(firstDone) + req := httptest.NewRequest(http.MethodGet, "/stream/direct/"+token, nil) + req.Header.Set("Range", "bytes=0-32767") + server.Handler().ServeHTTP(firstWriter, req) + }() + select { + case <-firstWriter.entered: + case <-time.After(2 * time.Second): + t.Fatal("first mounted range request did not begin its pour") + } + + second := httptest.NewRequest(http.MethodGet, "/stream/direct/"+token, nil) + second.Header.Set("Range", "bytes=32768-65535") + secondRec := httptest.NewRecorder() + server.Handler().ServeHTTP(secondRec, second) + if secondRec.Code != http.StatusPartialContent { + t.Fatalf("second range status = %d, want 206", secondRec.Code) + } + snapshot := server.tracker.Snapshot() + if len(snapshot) != 1 || snapshot[0].BytesServed == 0 { + t.Fatalf("first request lost tracking after second completed: %+v", snapshot) + } + key := nodesessions.KeyPrefix + server.tracker.NodeHash() + ":session-1" + if !redisState.has(key) { + t.Fatal("Redis key was deleted while the first range was still pouring") + } + + close(firstWriter.release) + select { + case <-firstDone: + case <-time.After(2 * time.Second): + t.Fatal("first range request did not complete") + } + if got := server.tracker.Snapshot(); len(got) != 0 { + t.Fatalf("final request release left session behind: %+v", got) + } +} + +func TestMountedDirectReleasesLeaseOnEveryCompletionPath(t *testing.T) { + path := t.TempDir() + "/movie.bin" + if err := os.WriteFile(path, []byte("payload"), 0o600); err != nil { + t.Fatal(err) + } + server, _, token := newMountedProxyTestServer(t, path) + + t.Run("head", func(t *testing.T) { + rec := httptest.NewRecorder() + server.Handler().ServeHTTP(rec, httptest.NewRequest(http.MethodHead, "/stream/direct/"+token, nil)) + if rec.Code != http.StatusOK { + t.Fatalf("status = %d", rec.Code) + } + if got := server.tracker.Snapshot(); len(got) != 0 { + t.Fatalf("HEAD leaked lease: %+v", got) + } + }) + + t.Run("missing file", func(t *testing.T) { + missingServer, _, missingToken := newMountedProxyTestServer(t, path+".missing") + rec := httptest.NewRecorder() + missingServer.Handler().ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/stream/direct/"+missingToken, nil)) + if rec.Code != http.StatusNotFound { + t.Fatalf("status = %d", rec.Code) + } + if got := missingServer.tracker.Snapshot(); len(got) != 0 { + t.Fatalf("missing-file completion leaked lease: %+v", got) + } + }) + + t.Run("client cancellation", func(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + req := httptest.NewRequest(http.MethodGet, "/stream/direct/"+token, nil).WithContext(ctx) + cancel() + rec := httptest.NewRecorder() + server.Handler().ServeHTTP(rec, req) + if got := server.tracker.Snapshot(); len(got) != 0 { + t.Fatalf("canceled request leaked lease: %+v", got) + } + }) +} + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +type limitedFailureWriter struct { + header http.Header + remaining int +} + +func (w *limitedFailureWriter) Header() http.Header { return w.header } +func (w *limitedFailureWriter) WriteHeader(int) {} +func (w *limitedFailureWriter) Write(p []byte) (int, error) { + if w.remaining <= 0 { + return 0, errors.New("client disconnected") + } + n := min(len(p), w.remaining) + w.remaining -= n + return n, errors.New("client disconnected") +} + +func TestMountedTranscodeLivenessRequiresSuccessfulServedBytes(t *testing.T) { + tests := []struct { + name string + method string + status int + body string + transportErr error + writer http.ResponseWriter + wantBytes int64 + wantServed bool + }{ + {name: "200 bytes", method: http.MethodGet, status: http.StatusOK, body: "manifest", wantBytes: 8, wantServed: true}, + {name: "206 bytes", method: http.MethodGet, status: http.StatusPartialContent, body: "segment", wantBytes: 7, wantServed: true}, + {name: "head zero bytes", method: http.MethodHead, status: http.StatusOK}, + {name: "404 error body", method: http.MethodGet, status: http.StatusNotFound, body: "not found"}, + {name: "upstream failure", method: http.MethodGet, transportErr: errors.New("dial failed")}, + {name: "client write failure", method: http.MethodGet, status: http.StatusOK, body: "abcdefgh", writer: &limitedFailureWriter{header: make(http.Header), remaining: 3}, wantBytes: 3, wantServed: true}, + } + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + server, _, token := newMountedProxyTestServer(t, "") + server.httpClient = &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { + if tc.transportErr != nil { + return nil, tc.transportErr + } + return &http.Response{ + StatusCode: tc.status, + Header: make(http.Header), + Body: io.NopCloser(bytes.NewBufferString(tc.body)), + }, nil + })} + writer := tc.writer + if writer == nil { + writer = httptest.NewRecorder() + } + server.Handler().ServeHTTP(writer, httptest.NewRequest(tc.method, "/stream/transcode/"+token+"/master.m3u8", nil)) + + got := server.tracker.Snapshot() + if len(got) != 1 { + t.Fatalf("snapshot = %+v, want visible ephemeral record", got) + } + if got[0].BytesServed != tc.wantBytes { + t.Fatalf("BytesServed = %d, want %d", got[0].BytesServed, tc.wantBytes) + } + if (got[0].LastServedAt != "") != tc.wantServed { + t.Fatalf("LastServedAt = %q, want served=%v", got[0].LastServedAt, tc.wantServed) + } + }) + } +} diff --git a/internal/streamenforcer/enforcer_test.go b/internal/streamenforcer/enforcer_test.go index cf1f8789..8397f6b2 100644 --- a/internal/streamenforcer/enforcer_test.go +++ b/internal/streamenforcer/enforcer_test.go @@ -7,6 +7,7 @@ import ( "testing" "time" + "github.com/Silo-Server/silo-server/internal/nodesessions" "github.com/Silo-Server/silo-server/internal/streammonitor" ) @@ -136,3 +137,20 @@ func TestEvaluateOnceSourceError(t *testing.T) { t.Fatalf("expected no revokes on source error, got %v", rev.revoked) } } + +func TestEvaluateOnceRevokesCanonicalLogicalSessionID(t *testing.T) { + source := streammonitor.NewFuncSource(func(context.Context) ([]nodesessions.SessionInfo, error) { + return []nodesessions.SessionInfo{ + {SessionID: "transport-a", LogicalSessionID: "logical-a", AuthUserID: 7, LastServedAt: "2026-07-30T01:00:00Z"}, + {SessionID: "transport-b", LogicalSessionID: "logical-b", AuthUserID: 7, LastServedAt: "2026-07-30T02:00:00Z"}, + }, nil + }) + rev := &fakeRevoker{} + e := New(source, func(context.Context, int) (int, error) { return 1, nil }, rev, 0) + if err := e.EvaluateOnce(context.Background()); err != nil { + t.Fatalf("EvaluateOnce: %v", err) + } + if len(rev.revoked) != 1 || rev.revoked[0] != "logical-a" { + t.Fatalf("revoked = %v, want logical-a", rev.revoked) + } +} diff --git a/internal/streammonitor/monitor.go b/internal/streammonitor/monitor.go index 64e14640..0c633ef4 100644 --- a/internal/streammonitor/monitor.go +++ b/internal/streammonitor/monitor.go @@ -29,6 +29,7 @@ import ( "github.com/redis/go-redis/v9" "github.com/Silo-Server/silo-server/internal/nodesessions" + "github.com/Silo-Server/silo-server/internal/playback" ) // sessionKeyPrefix is the Redis key prefix under which nodesessions.Tracker @@ -42,21 +43,22 @@ const scanCount = 256 // LiveStream is a normalized view of a single active streaming session. type LiveStream struct { - SessionID string - UserID int // from SessionInfo.AuthUserID - ProfileID string - NodeName string - NodeURL string - Type string // play method: direct_play | remux | transcode - Route string // origin protocol: native | jellycompat - MediaFileID int - ClientIP string - ClientName string - Position float64 // last known playback position (seconds); secondary timing - HWAccel string - LastServedAt time.Time // parsed from SessionInfo.LastServedAt (zero if absent) - BytesServed int64 - StartedAt time.Time // parsed from SessionInfo.StartedAt (zero if unparseable) + SessionID string + LogicalSessionID string + UserID int // from SessionInfo.AuthUserID + ProfileID string + NodeName string + NodeURL string + Type string // play method: direct_play | remux | transcode + Route string // origin protocol: native | jellycompat + MediaFileID int + ClientIP string + ClientName string + Position float64 // last known playback position (seconds); secondary timing + HWAccel string + LastServedAt time.Time // parsed from SessionInfo.LastServedAt (zero if absent) + BytesServed int64 + StartedAt time.Time // parsed from SessionInfo.StartedAt (zero if unparseable) } // Snapshot is a point-in-time picture of the live streams. @@ -117,7 +119,8 @@ func NewMultiSource(sources ...Source) *MultiSource { return &MultiSource{sources: sources} } -// Snapshot merges every sub-source's snapshot, de-duplicated by session id. +// Snapshot merges every sub-source's snapshot, de-duplicated by canonical +// logical session identity when available. func (m *MultiSource) Snapshot(ctx context.Context) (Snapshot, error) { var all []LiveStream for _, src := range m.sources { @@ -139,24 +142,35 @@ func (m *MultiSource) Snapshot(ctx context.Context) (Snapshot, error) { // the zero time). func toLiveStream(info nodesessions.SessionInfo) LiveStream { return LiveStream{ - SessionID: info.SessionID, - UserID: info.AuthUserID, - ProfileID: info.ProfileID, - NodeName: info.NodeName, - NodeURL: info.NodeURL, - Type: info.Type, - Route: info.Route, - MediaFileID: info.MediaFileID, - ClientIP: info.ClientIP, - ClientName: info.ClientName, - Position: info.Position, - HWAccel: info.HWAccel, - LastServedAt: parseTime(info.LastServedAt), - BytesServed: info.BytesServed, - StartedAt: parseTime(info.StartedAt), + SessionID: info.SessionID, + LogicalSessionID: info.LogicalSessionID, + UserID: info.AuthUserID, + ProfileID: info.ProfileID, + NodeName: info.NodeName, + NodeURL: info.NodeURL, + Type: info.Type, + Route: info.Route, + MediaFileID: info.MediaFileID, + ClientIP: info.ClientIP, + ClientName: info.ClientName, + Position: info.Position, + HWAccel: info.HWAccel, + LastServedAt: parseTime(info.LastServedAt), + BytesServed: info.BytesServed, + StartedAt: parseTime(info.StartedAt), } } +// sessionIdentity returns the stable identity used to merge and enforce a +// playback stream. SessionID remains the transport-addressing key on raw node +// records; a logical id, when present, takes precedence only in monitoring. +func sessionIdentity(sessionID, logicalSessionID string) string { + if logicalSessionID != "" { + return logicalSessionID + } + return sessionID +} + // parseTime parses an RFC3339 timestamp, returning the zero time for empty or // unparseable input. func parseTime(s string) time.Time { @@ -171,11 +185,11 @@ func parseTime(s string) time.Time { return t } -// mergeStreams collapses records that share a SessionID (the same session can be -// tracked by more than one node — e.g. a proxy record and a transcode-node -// record), keeping the one with the most recent LastServedAt. Records with a -// distinct SessionID are all retained. The relative order of the kept records is -// not guaranteed. +// mergeStreams collapses records that share a canonical identity (the logical +// session id when present, otherwise the transport SessionID). The same stream +// can be tracked by more than one node or transport generation. The freshest +// record wins; genuinely distinct logical sessions are retained. The relative +// order of kept records is not guaranteed. // // Ownership and attribution are carried forward independently of the freshness // pick: the transcode node's own start record has no resolved owner (UserID 0) @@ -188,9 +202,11 @@ func parseTime(s string) time.Time { func mergeStreams(streams []LiveStream) []LiveStream { bySession := make(map[string]LiveStream, len(streams)) for _, st := range streams { - existing, ok := bySession[st.SessionID] + key := sessionIdentity(st.SessionID, st.LogicalSessionID) + existing, ok := bySession[key] if !ok { - bySession[st.SessionID] = st + st.SessionID = key + bySession[key] = st continue } winner := existing @@ -229,7 +245,8 @@ func mergeStreams(streams []LiveStream) []LiveStream { } // These are observers of one pour, not independent byte sources. winner.BytesServed = max(winner.BytesServed, other.BytesServed) - bySession[st.SessionID] = winner + winner.SessionID = key + bySession[key] = winner } out := make([]LiveStream, 0, len(bySession)) for _, st := range bySession { @@ -238,22 +255,21 @@ func mergeStreams(streams []LiveStream) []LiveStream { return out } -// DedupeSessionInfos collapses raw monitoring records that share a SessionID, -// applying the same rules as mergeStreams — keep the most-recently-served copy, -// carry a resolved owner forward, backfill missing attribution — but preserving -// the SessionInfo shape for surfaces whose wire format IS the raw record (the -// admin session list, which unions Redis edge records with the in-process -// integrated sessions and would otherwise show the same stream twice). Kept as -// a sibling of mergeStreams rather than a shared generic because mergeStreams -// operates on the parsed LiveStream form; keep the two rule sets in sync. +// DedupeSessionInfos collapses raw monitoring records that share a canonical +// identity, applying the same rules as mergeStreams — keep the most-recently- +// served copy, carry a resolved owner forward, backfill missing attribution — +// but preserving the SessionInfo shape and transport SessionID for surfaces +// whose wire format IS the raw record. Kept as a sibling of mergeStreams rather +// than a shared generic because mergeStreams operates on parsed LiveStream. func DedupeSessionInfos(infos []nodesessions.SessionInfo) []nodesessions.SessionInfo { bySession := make(map[string]nodesessions.SessionInfo, len(infos)) order := make([]string, 0, len(infos)) for _, in := range infos { - existing, ok := bySession[in.SessionID] + key := sessionIdentity(in.SessionID, in.LogicalSessionID) + existing, ok := bySession[key] if !ok { - bySession[in.SessionID] = in - order = append(order, in.SessionID) + bySession[key] = in + order = append(order, key) continue } winner, other := existing, in @@ -281,7 +297,7 @@ func DedupeSessionInfos(infos []nodesessions.SessionInfo) []nodesessions.Session winner.Position = other.Position } winner.BytesServed = max(winner.BytesServed, other.BytesServed) - bySession[in.SessionID] = winner + bySession[key] = winner } out := make([]nodesessions.SessionInfo, 0, len(bySession)) for _, id := range order { @@ -290,6 +306,40 @@ func DedupeSessionInfos(infos []nodesessions.SessionInfo) []nodesessions.Session return out } +// LiveLocalSessions maps in-process playback sessions into monitoring records +// so integrated streams are visible to both the enforcer and admin view. +func LiveLocalSessions(sm *playback.SessionManager, nodeName string) []nodesessions.SessionInfo { + if sm == nil { + return nil + } + live := sm.AllSessions() + out := make([]nodesessions.SessionInfo, 0, len(live)) + for _, s := range live { + lastServedAt := s.LastServedAt + if lastServedAt.IsZero() { + lastServedAt = s.LastActivityAt + } + out = append(out, nodesessions.SessionInfo{ + SessionID: s.ID, + NodeName: nodeName, + AuthUserID: s.UserID, + ProfileID: s.ProfileID, + Type: string(s.PlayMethod), + Route: s.Origin(), + MediaFileID: s.MediaFileID, + ClientIP: s.ClientIP, + ClientName: s.ClientName, + Position: s.Position, + Resolution: s.TargetResolution, + HWAccel: s.TranscodeHWAccel, + StartedAt: s.StartedAt.UTC().Format(time.RFC3339), + LastServedAt: lastServedAt.UTC().Format(time.RFC3339), + BytesServed: s.BytesServed, + }) + } + return out +} + // RedisSource reads silo:sessions:* records, producing the multi-node // authoritative picture. type RedisSource struct { @@ -302,8 +352,8 @@ func NewRedisSource(rdb *redis.Client) *RedisSource { } // Snapshot SCANs every silo:sessions:* record, decodes it, and returns the -// deduped live picture. The same session appearing on multiple nodes is -// collapsed to the record with the most recent LastServedAt. +// deduped live picture. The same logical session appearing on multiple nodes or +// transport generations is collapsed to the record with the newest liveness. func (r *RedisSource) Snapshot(ctx context.Context) (Snapshot, error) { if r.rdb == nil { return Snapshot{Streams: []LiveStream{}}, nil diff --git a/internal/streammonitor/monitor_test.go b/internal/streammonitor/monitor_test.go index 8ae09a29..31dd9a77 100644 --- a/internal/streammonitor/monitor_test.go +++ b/internal/streammonitor/monitor_test.go @@ -6,6 +6,7 @@ import ( "time" "github.com/Silo-Server/silo-server/internal/nodesessions" + "github.com/Silo-Server/silo-server/internal/playback" ) func fakeFn(infos []nodesessions.SessionInfo) func(ctx context.Context) ([]nodesessions.SessionInfo, error) { @@ -318,3 +319,63 @@ func TestDedupeSessionInfosKeepsLargestObservedByteTotal(t *testing.T) { }) } } + +func TestLogicalAndTransportRecordsMergeWithOwnerResolved(t *testing.T) { + out := mergeStreams([]LiveStream{ + {SessionID: "logical", UserID: 17, ProfileID: "p", LastServedAt: time.Unix(100, 0)}, + {SessionID: "transport-a", LogicalSessionID: "logical", LastServedAt: time.Unix(200, 0)}, + }) + if len(out) != 1 { + t.Fatalf("merge len = %d, want 1", len(out)) + } + if out[0].SessionID != "logical" || out[0].UserID != 17 || out[0].ProfileID != "p" { + t.Fatalf("canonical merged stream = %+v", out[0]) + } +} + +func TestDistinctLogicalStreamsStayDistinct(t *testing.T) { + out := mergeStreams([]LiveStream{ + {SessionID: "transport-a", LogicalSessionID: "logical-a"}, + {SessionID: "transport-b", LogicalSessionID: "logical-b"}, + }) + if len(out) != 2 { + t.Fatalf("distinct logical streams merged: %+v", out) + } +} + +func TestTransportGenerationsOfOneLogicalSessionMerge(t *testing.T) { + out := DedupeSessionInfos([]nodesessions.SessionInfo{ + {SessionID: "transport-a", LogicalSessionID: "logical", NodeName: "old", LastServedAt: "2026-07-30T01:00:00Z"}, + {SessionID: "transport-b", LogicalSessionID: "logical", NodeName: "new", LastServedAt: "2026-07-30T02:00:00Z"}, + }) + if len(out) != 1 || out[0].SessionID != "transport-b" || out[0].LogicalSessionID != "logical" { + t.Fatalf("transport generation dedupe = %+v", out) + } +} + +func TestLiveLocalSessionsMapping(t *testing.T) { + sm := playback.NewSessionManager(0, 0) + ctx := playback.WithClientInfo(context.Background(), playback.ClientInfo{ + Name: "Silo TV", + }) + session, err := sm.StartSessionWithContext(ctx, 42, "profile-1", 9, playback.PlayDirect, false) + if err != nil { + t.Fatalf("StartSessionWithContext: %v", err) + } + session.ClientIP = "192.0.2.10" + got := LiveLocalSessions(sm, "local") + if len(got) != 1 { + t.Fatalf("sessions = %+v", got) + } + info := got[0] + if info.SessionID != session.ID || info.NodeName != "local" || + info.AuthUserID != 42 || info.ProfileID != "profile-1" || + info.MediaFileID != 9 || info.Type != string(playback.PlayDirect) || + info.Route != session.Origin() || info.ClientName != "Silo TV" || + info.ClientIP != "192.0.2.10" { + t.Fatalf("mapped session = %+v", info) + } + if info.LastServedAt != session.LastActivityAt.UTC().Format(time.RFC3339) { + t.Fatalf("LastServedAt = %q, want LastActivityAt fallback", info.LastServedAt) + } +} diff --git a/internal/transcodenode/server.go b/internal/transcodenode/server.go index d2920817..df566896 100644 --- a/internal/transcodenode/server.go +++ b/internal/transcodenode/server.go @@ -28,6 +28,7 @@ import ( // TranscodeStartRequest is the JSON body for POST /transcode/start. type TranscodeStartRequest struct { SessionID string `json:"session_id"` + LogicalSessionID string `json:"logical_session_id,omitempty"` InputPath string `json:"input_path"` SourceVideoCodec string `json:"source_video_codec"` VideoBitstreamFilter string `json:"video_bitstream_filter,omitempty"` @@ -119,6 +120,11 @@ type Server struct { lifecycleMu sync.Mutex lifecycleLocks map[string]*sessionLifecycleLock + // runTracker is a deterministic scheduling seam for lifecycle tests. + // Production leaves it nil and launches each write directly in a goroutine; + // this is not a projection queue. + runTracker func(func()) + // recipeStore is the control-plane recipe store consulted when a forwarded // token carries no recipe (the jellycompat node hop). Nil disables that path. recipeStore recipeStore @@ -160,6 +166,29 @@ func (s *Server) lockSessionLifecycle(sessionID string) func() { } } +// trackIfCurrent serializes the delayed monitoring write with stop/reap/start +// lifecycle transitions. Once it owns the lifecycle lock, it verifies that the +// exact session generation is still registered before writing; a stopped or +// replaced generation can therefore never recreate a ghost record. +func (s *Server) trackIfCurrent(ctx context.Context, sessionID string, expected *playback.TranscodeSession, info nodesessions.SessionInfo) { + run := func() { + unlock := s.lockSessionLifecycle(sessionID) + defer unlock() + s.mu.RLock() + current, ok := s.sessions[sessionID] + s.mu.RUnlock() + if !ok || current != expected { + return + } + s.tracker.Track(ctx, info) + } + if s.runTracker != nil { + s.runTracker(run) + return + } + go run() +} + // restartSessionLocked re-spawns session under the per-session lifecycle lock so // a segment-recovery restart can never race a fresh start, reconstruct, or // another restart into the same output directory. It holds the lock only across @@ -601,21 +630,22 @@ func (s *Server) handleStart(w http.ResponseWriter, r *http.Request) { // tracking write is monitoring-only. effectiveHWAccel := session.Opts().HWAccel trackCtx := context.WithoutCancel(r.Context()) - go s.tracker.Track(trackCtx, nodesessions.SessionInfo{ - SessionID: req.SessionID, - NodeURL: s.tracker.NodeURL(), - NodeName: s.tracker.NodeName(), - Type: "transcode", - CodecVideo: req.TargetCodecVideo, - CodecAudio: req.TargetCodecAudio, - Resolution: req.TargetResolution, - HWAccel: effectiveHWAccel, - StartedAt: time.Now().UTC().Format(time.RFC3339), - AuthUserID: req.AuthUserID, - ProfileID: req.ProfileID, - MediaFileID: req.MediaFileID, - Route: req.Route, - ClientName: req.ClientName, + s.trackIfCurrent(trackCtx, req.SessionID, session, nodesessions.SessionInfo{ + SessionID: req.SessionID, + LogicalSessionID: req.LogicalSessionID, + NodeURL: s.tracker.NodeURL(), + NodeName: s.tracker.NodeName(), + Type: "transcode", + CodecVideo: req.TargetCodecVideo, + CodecAudio: req.TargetCodecAudio, + Resolution: req.TargetResolution, + HWAccel: effectiveHWAccel, + StartedAt: time.Now().UTC().Format(time.RFC3339), + AuthUserID: req.AuthUserID, + ProfileID: req.ProfileID, + MediaFileID: req.MediaFileID, + Route: req.Route, + ClientName: req.ClientName, }) w.WriteHeader(http.StatusAccepted) @@ -781,21 +811,22 @@ func (s *Server) spawnReconstruct(r *http.Request, sessionID string, requestedSe s.activeJobs.Add(1) trackCtx := context.WithoutCancel(r.Context()) - go s.tracker.Track(trackCtx, nodesessions.SessionInfo{ - SessionID: sessionID, - NodeURL: s.tracker.NodeURL(), - NodeName: s.tracker.NodeName(), - Type: "transcode", - CodecVideo: card.TargetCodecVideo, - CodecAudio: card.TargetCodecAudio, - Resolution: card.TargetResolution, - HWAccel: session.Opts().HWAccel, - StartedAt: time.Now().UTC().Format(time.RFC3339), - AuthUserID: card.UserID, - ProfileID: card.ProfileID, - MediaFileID: card.MediaFileID, - Route: route, - ClientName: clientName, + s.trackIfCurrent(trackCtx, sessionID, session, nodesessions.SessionInfo{ + SessionID: sessionID, + LogicalSessionID: card.SessionID, + NodeURL: s.tracker.NodeURL(), + NodeName: s.tracker.NodeName(), + Type: "transcode", + CodecVideo: card.TargetCodecVideo, + CodecAudio: card.TargetCodecAudio, + Resolution: card.TargetResolution, + HWAccel: session.Opts().HWAccel, + StartedAt: time.Now().UTC().Format(time.RFC3339), + AuthUserID: card.UserID, + ProfileID: card.ProfileID, + MediaFileID: card.MediaFileID, + Route: route, + ClientName: clientName, }) slog.InfoContext(r.Context(), "transcode node session reconstructed from token", "component", "transcodenode", diff --git a/internal/transcodenode/tracker_lifecycle_test.go b/internal/transcodenode/tracker_lifecycle_test.go new file mode 100644 index 00000000..eabc6a70 --- /dev/null +++ b/internal/transcodenode/tracker_lifecycle_test.go @@ -0,0 +1,129 @@ +package transcodenode + +import ( + "context" + "net" + "sync" + "testing" + + "github.com/redis/go-redis/v9" + + "github.com/Silo-Server/silo-server/internal/nodesessions" + "github.com/Silo-Server/silo-server/internal/playback" +) + +type transcodeRedisHook struct { + mu sync.Mutex + keys map[string]bool +} + +func (h *transcodeRedisHook) DialHook(next redis.DialHook) redis.DialHook { return next } +func (h *transcodeRedisHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { + return func(ctx context.Context, cmd redis.Cmder) error { + h.mu.Lock() + defer h.mu.Unlock() + switch cmd.Name() { + case "set": + h.keys[cmd.Args()[1].(string)] = true + if status, ok := cmd.(*redis.StatusCmd); ok { + status.SetVal("OK") + } + return nil + case "del": + for _, arg := range cmd.Args()[1:] { + delete(h.keys, arg.(string)) + } + if count, ok := cmd.(*redis.IntCmd); ok { + count.SetVal(1) + } + return nil + } + return next(ctx, cmd) + } +} +func (h *transcodeRedisHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return next +} + +func newTranscodeLifecycleTracker(t *testing.T) (*nodesessions.Tracker, *transcodeRedisHook) { + t.Helper() + rdb := redis.NewClient(&redis.Options{Dialer: func(context.Context, string, string) (net.Conn, error) { + t.Fatal("unexpected Redis dial") + return nil, nil + }}) + t.Cleanup(func() { _ = rdb.Close() }) + hook := &transcodeRedisHook{keys: make(map[string]bool)} + rdb.AddHook(hook) + return nodesessions.NewTracker(rdb, "http://node", "node", "transcode"), hook +} + +func (h *transcodeRedisHook) has(key string) bool { + h.mu.Lock() + defer h.mu.Unlock() + return h.keys[key] +} + +func TestDelayedTrackSkipsStoppedSession(t *testing.T) { + tracker, redisState := newTranscodeLifecycleTracker(t) + session := &playback.TranscodeSession{} + var queued []func() + s := &Server{ + tracker: tracker, + sessions: map[string]*playback.TranscodeSession{"transport": session}, + runTracker: func(fn func()) { + queued = append(queued, fn) + }, + } + s.trackIfCurrent(context.Background(), "transport", session, nodesessions.SessionInfo{ + SessionID: "transport", LogicalSessionID: "logical", + }) + + unlock := s.lockSessionLifecycle("transport") + s.mu.Lock() + delete(s.sessions, "transport") + s.mu.Unlock() + tracker.Remove(context.Background(), "transport") + unlock() + queued[0]() + + if got := tracker.Snapshot(); len(got) != 0 { + t.Fatalf("delayed track recreated stopped record: %+v", got) + } + key := nodesessions.KeyPrefix + tracker.NodeHash() + ":transport" + if redisState.has(key) { + t.Fatal("delayed track recreated stopped Redis key") + } +} + +func TestDelayedTrackSkipsReplacedSession(t *testing.T) { + tracker, _ := newTranscodeLifecycleTracker(t) + oldSession := &playback.TranscodeSession{} + newSession := &playback.TranscodeSession{} + var queued []func() + s := &Server{ + tracker: tracker, + sessions: map[string]*playback.TranscodeSession{"transport": oldSession}, + runTracker: func(fn func()) { + queued = append(queued, fn) + }, + } + s.trackIfCurrent(context.Background(), "transport", oldSession, nodesessions.SessionInfo{ + SessionID: "transport", LogicalSessionID: "old", + }) + + unlock := s.lockSessionLifecycle("transport") + s.mu.Lock() + s.sessions["transport"] = newSession + s.mu.Unlock() + tracker.Remove(context.Background(), "transport") + tracker.Track(context.Background(), nodesessions.SessionInfo{ + SessionID: "transport", LogicalSessionID: "new", + }) + unlock() + queued[0]() + + got := tracker.Snapshot() + if len(got) != 1 || got[0].LogicalSessionID != "new" { + t.Fatalf("old delayed track replaced new record: %+v", got) + } +}