diff --git a/pkg/api/api.go b/pkg/api/api.go index 1c4e8760d..64295c80b 100644 --- a/pkg/api/api.go +++ b/pkg/api/api.go @@ -77,6 +77,8 @@ type StreamplaceAPI struct { // override tls port for http redirect server if we're using systemd file descriptors HTTPRedirectTLSPort *int + sessions map[string]map[string]time.Time + sessionsLock sync.RWMutex } type WebsocketTracker struct { @@ -106,6 +108,8 @@ func MakeStreamplaceAPI(cli *config.CLI, mod model.Model, statefulDB *statedb.St limiters: make(map[string]*rate.Limiter), SignerCache: make(map[string]media.MediaSigner), op: op, + sessions: make(map[string]map[string]time.Time), + sessionsLock: sync.RWMutex{}, } a.Mimes, err = updater.GetMimes() if err != nil { diff --git a/pkg/api/playback.go b/pkg/api/playback.go index 73ca2c1d5..b48941717 100644 --- a/pkg/api/playback.go +++ b/pkg/api/playback.go @@ -177,6 +177,46 @@ func NoCache(h httprouter.Handle) httprouter.Handle { } } +const SessionExpireTime = 30 * time.Second + +func (a *StreamplaceAPI) SessionSeen(ctx context.Context, user string, session string) { + now := time.Now() + go func() { + a.sessionsLock.Lock() + defer a.sessionsLock.Unlock() + if _, ok := a.sessions[user]; !ok { + a.sessions[user] = map[string]time.Time{} + } + if _, ok := a.sessions[user][session]; !ok { + log.Warn(ctx, "ViewerInc", "user", user, "session", session) + spmetrics.ViewerInc(user, "hls") + a.Bus.IncrementViewerCount(user, "local") + } + a.sessions[user][session] = now + }() +} + +func (a *StreamplaceAPI) ExpireSessions(ctx context.Context) error { + for { + select { + case <-ctx.Done(): + return nil + case <-time.After(5 * time.Second): + a.sessionsLock.Lock() + for user, sessions := range a.sessions { + for session, seen := range sessions { + if time.Since(seen) > SessionExpireTime { + delete(sessions, session) + spmetrics.ViewerDec(user, "hls") + a.Bus.DecrementViewerCount(user, "local") + } + } + } + a.sessionsLock.Unlock() + } + } +} + func (a *StreamplaceAPI) HandleHLSPlayback(ctx context.Context) httprouter.Handle { return NoCache(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) { user := p.ByName("user") @@ -211,7 +251,7 @@ func (a *StreamplaceAPI) HandleHLSPlayback(ctx context.Context) httprouter.Handl w.Header().Set("Content-Type", "application/vnd.apple.mpegurl") } else { if session != "" { - spmetrics.SessionSeen(user, session) + a.SessionSeen(ctx, user, session) } w.Header().Set("Content-Type", "video/mp2t") } diff --git a/pkg/api/websocket.go b/pkg/api/websocket.go index c11d7797f..96d9e79ca 100644 --- a/pkg/api/websocket.go +++ b/pkg/api/websocket.go @@ -139,7 +139,7 @@ func (a *StreamplaceAPI) HandleWebsocket(ctx context.Context) httprouter.Handle case msg := <-initialBurst: send(msg) case <-ticker.C: - count := spmetrics.GetViewCount(repoDID) + count := a.Bus.GetViewerCount(repoDID) bs, err := json.Marshal(streamplace.Livestream_ViewerCount{Count: int64(count), LexiconTypeID: "place.stream.livestream#viewerCount"}) if err != nil { log.Error(ctx, "could not marshal view count", "error", err) @@ -226,7 +226,7 @@ func (a *StreamplaceAPI) HandleWebsocket(ctx context.Context) httprouter.Handle }() go func() { - count := spmetrics.GetViewCount(repoDID) + count := a.Bus.GetViewerCount(repoDID) initialBurst <- streamplace.Livestream_ViewerCount{Count: int64(count), LexiconTypeID: "place.stream.livestream#viewerCount"} }() diff --git a/pkg/bus/bus.go b/pkg/bus/bus.go index aa924933a..9f004223b 100644 --- a/pkg/bus/bus.go +++ b/pkg/bus/bus.go @@ -7,21 +7,32 @@ import ( type Message any type Subscription chan Message +type ViewerCountUpdate struct { + Streamer string + Count int + Origin string +} + // Bus is a simple pub/sub system for backing websocket connections type Bus struct { - mu sync.Mutex - clients map[string][]Subscription - segChans map[string][]*SegChan - segChansMutex sync.Mutex - segBuf map[string][]*Seg - segBufMutex sync.RWMutex + mu sync.Mutex + clients map[string][]Subscription + segChans map[string][]*SegChan + segChansMutex sync.Mutex + segBuf map[string][]*Seg + segBufMutex sync.RWMutex + viewerCounts map[string]map[string]int + viewerCountsMutex sync.RWMutex + viewerCountSubscriptions []chan ViewerCountUpdate } func NewBus() *Bus { return &Bus{ - clients: make(map[string][]Subscription), - segChans: make(map[string][]*SegChan), - segBuf: make(map[string][]*Seg), + clients: make(map[string][]Subscription), + segChans: make(map[string][]*SegChan), + segBuf: make(map[string][]*Seg), + viewerCounts: make(map[string]map[string]int), + viewerCountSubscriptions: []chan ViewerCountUpdate{}, } } @@ -59,6 +70,14 @@ func (b *Bus) Unsubscribe(user string, ch <-chan Message) { } } +func (b *Bus) SubscribeToViewerCount() <-chan ViewerCountUpdate { + b.viewerCountsMutex.Lock() + defer b.viewerCountsMutex.Unlock() + ch := make(chan ViewerCountUpdate, 100) + b.viewerCountSubscriptions = append(b.viewerCountSubscriptions, ch) + return ch +} + func (b *Bus) Publish(user string, msg Message) { b.mu.Lock() defer b.mu.Unlock() @@ -72,3 +91,59 @@ func (b *Bus) Publish(user string, msg Message) { }(sub) } } + +func (b *Bus) GetViewerCount(user string) int { + b.viewerCountsMutex.RLock() + defer b.viewerCountsMutex.RUnlock() + streamerCounts, ok := b.viewerCounts[user] + if !ok { + return 0 + } + count := 0 + for _, viewers := range streamerCounts { + count += viewers + } + return count +} + +func (b *Bus) SetViewerCount(user string, origin string, count int) { + b.viewerCountsMutex.Lock() + defer b.viewerCountsMutex.Unlock() + _, ok := b.viewerCounts[user] + if !ok { + b.viewerCounts[user] = make(map[string]int) + } + b.viewerCounts[user][origin] = count + b.notifyViewerCountSubscribers(user, count, origin) +} + +func (b *Bus) IncrementViewerCount(user string, origin string) { + b.viewerCountsMutex.Lock() + defer b.viewerCountsMutex.Unlock() + _, ok := b.viewerCounts[user] + if !ok { + b.viewerCounts[user] = make(map[string]int) + } + b.viewerCounts[user][origin] += 1 + b.notifyViewerCountSubscribers(user, b.viewerCounts[user][origin], origin) +} + +func (b *Bus) DecrementViewerCount(user string, origin string) { + b.viewerCountsMutex.Lock() + defer b.viewerCountsMutex.Unlock() + _, ok := b.viewerCounts[user] + if !ok { + b.viewerCounts[user] = make(map[string]int) + } + b.viewerCounts[user][origin] -= 1 + b.notifyViewerCountSubscribers(user, b.viewerCounts[user][origin], origin) +} + +// only call if you're holding viewerCountsMutex +func (b *Bus) notifyViewerCountSubscribers(user string, count int, origin string) { + for _, sub := range b.viewerCountSubscriptions { + go func() { + sub <- ViewerCountUpdate{Streamer: user, Count: count, Origin: origin} + }() + } +} diff --git a/pkg/cmd/streamplace.go b/pkg/cmd/streamplace.go index 30e2c0a12..46ba8ff46 100644 --- a/pkg/cmd/streamplace.go +++ b/pkg/cmd/streamplace.go @@ -503,7 +503,7 @@ func start(build *config.BuildFlags, platformJobs []jobFunc) error { } group.Go(func() error { - return spmetrics.ExpireSessions(ctx) + return a.ExpireSessions(ctx) }) group.Go(func() error { diff --git a/pkg/media/viewers.go b/pkg/media/viewers.go new file mode 100644 index 000000000..d8a8b3948 --- /dev/null +++ b/pkg/media/viewers.go @@ -0,0 +1,13 @@ +package media + +import "stream.place/streamplace/pkg/spmetrics" + +func (mm *MediaManager) IncrementViewerCount(user string, protocol string) { + mm.bus.IncrementViewerCount(user, "local") + spmetrics.ViewerInc(user, protocol) +} + +func (mm *MediaManager) DecrementViewerCount(user string, protocol string) { + mm.bus.DecrementViewerCount(user, "local") + spmetrics.ViewerDec(user, protocol) +} diff --git a/pkg/media/webrtc_playback.go b/pkg/media/webrtc_playback.go index b5538c8f0..bd69e50bc 100644 --- a/pkg/media/webrtc_playback.go +++ b/pkg/media/webrtc_playback.go @@ -13,7 +13,6 @@ import ( "github.com/pion/webrtc/v4/pkg/media" "stream.place/streamplace/pkg/bus" "stream.place/streamplace/pkg/log" - "stream.place/streamplace/pkg/spmetrics" ) // we have a bug that prevents us from correctly probing video durations @@ -303,8 +302,8 @@ func (mm *MediaManager) WebRTCPlayback(ctx context.Context, user string, renditi if err != nil { log.Log(ctx, "failed to set pipeline state to null", "error", err) } - spmetrics.ViewerInc(user, "webrtc") - defer spmetrics.ViewerDec(user, "webrtc") + mm.IncrementViewerCount(user, "webrtc") + defer mm.DecrementViewerCount(user, "webrtc") go func() { rtcpBuf := make([]byte, 1500) diff --git a/pkg/media/webrtc_playback2.go b/pkg/media/webrtc_playback2.go index e567f6fb1..5ee005afb 100644 --- a/pkg/media/webrtc_playback2.go +++ b/pkg/media/webrtc_playback2.go @@ -11,7 +11,6 @@ import ( "golang.org/x/sync/errgroup" "stream.place/streamplace/pkg/bus" "stream.place/streamplace/pkg/log" - "stream.place/streamplace/pkg/spmetrics" ) // This function remains in scope for the duration of a single users' playback @@ -181,8 +180,8 @@ func (mm *MediaManager) WebRTCPlayback2(ctx context.Context, user string, rendit } }() - spmetrics.ViewerInc(user, "webrtc") - defer spmetrics.ViewerDec(user, "webrtc") + mm.IncrementViewerCount(user, "webrtc") + defer mm.DecrementViewerCount(user, "webrtc") go func() { rtcpBuf := make([]byte, 1500) diff --git a/pkg/replication/iroh_replicator/kv.go b/pkg/replication/iroh_replicator/kv.go index 251e3b8b2..5b3fd69d7 100644 --- a/pkg/replication/iroh_replicator/kv.go +++ b/pkg/replication/iroh_replicator/kv.go @@ -6,6 +6,8 @@ import ( "crypto/rand" "encoding/json" "fmt" + "reflect" + "strings" "sync" "time" @@ -28,7 +30,7 @@ type IrohSwarm struct { segChan chan *media.NewSegmentNotification NodeID string NodeTicket string - activeSubs map[string]*OriginInfo + activeSubs map[string]*SwarmOriginInfo handleDataScoped func(topic string, data []byte) bus *bus.Bus originMutex sync.Mutex @@ -37,9 +39,18 @@ type IrohSwarm struct { } // A message saying "hey I ingested node data at this time" -type OriginInfo struct { - NodeID string `json:"node_id"` - Time string `json:"time"` +type SwarmOriginInfo struct { + Type string `json:"$type"` + NodeID string `json:"node_id"` + Time string `json:"time"` + Streamer string `json:"streamer"` +} + +type SwarmViewerCount struct { + Type string `json:"$type"` + Server string `json:"server"` + Streamer string `json:"streamer"` + Viewers int `json:"viewers"` } func NewSwarm(ctx context.Context, cli *config.CLI, secret []byte, topic []byte, mm *media.MediaManager, bus *bus.Bus, mod model.Model) (*IrohSwarm, error) { @@ -63,7 +74,7 @@ func NewSwarm(ctx context.Context, cli *config.CLI, secret []byte, topic []byte, swarm := IrohSwarm{ mm: mm, - activeSubs: make(map[string]*OriginInfo), + activeSubs: make(map[string]*SwarmOriginInfo), bus: bus, mod: mod, cli: cli, @@ -137,6 +148,9 @@ func (swarm *IrohSwarm) Start(ctx context.Context, tickets []string) error { g.Go(func() error { return swarm.startBusSubscribe(ctx) }) + g.Go(func() error { + return swarm.startViewerCountSubscribe(ctx) + }) return g.Wait() } @@ -152,31 +166,19 @@ func (swarm *IrohSwarm) startKV(ctx context.Context) error { } if ev == nil { - log.Debug(ctx, "Got empty event from sub.NextRaw(), pausing for a second") + log.Warn(ctx, "Got empty event from sub.NextRaw(), pausing for a second and continuing") time.Sleep(1 * time.Second) continue } + switch item := (*ev).(type) { case iroh_streamplace.SubscribeItemEntry: - keyStr := string(item.Key) - valueStr := string(item.Value) - log.Debug(ctx, "SubscribeItemEntry", "key", keyStr, "value", valueStr) - if len(valueStr) > 0 && valueStr[0] != '{' { - // not JSON, it's one of the rust messages - log.Debug(ctx, "not JSON", "key", keyStr, "value", valueStr) - continue - } - var info OriginInfo - err := json.Unmarshal(item.Value, &info) + err := swarm.handleIrohMessage(ctx, item) if err != nil { - log.Error(ctx, "could not unmarshal origin info", "error", err) - continue - } - err = swarm.checkOrigins(ctx, keyStr, info.NodeID) - if err != nil { - log.Error(ctx, "could not check origins", "error", err) + log.Error(ctx, "could not handle iroh message", "error", err) continue } + case iroh_streamplace.SubscribeItemCurrentDone: log.Debug(ctx, "SubscribeItemCurrentDone", "currentDone", item) case iroh_streamplace.SubscribeItemExpired: @@ -187,6 +189,61 @@ func (swarm *IrohSwarm) startKV(ctx context.Context) error { } } +func (swarm *IrohSwarm) handleIrohMessage(ctx context.Context, item iroh_streamplace.SubscribeItemEntry) error { + keyStr := string(item.Key) + valueStr := string(item.Value) + log.Warn(ctx, "SubscribeItemEntry", "key", keyStr, "value", valueStr) + if len(valueStr) > 0 && valueStr[0] != '{' { + // not JSON, it's one of the rust messages + log.Debug(ctx, "not JSON", "key", keyStr, "value", valueStr) + return nil + } + rawMessage, err := decodeIrohMessage(item.Key, item.Value) + if err != nil { + return fmt.Errorf("could not decode iroh message: %w", err) + } + switch message := rawMessage.(type) { + case SwarmOriginInfo: + err = swarm.checkOrigins(ctx, message.Streamer, message.NodeID) + if err != nil { + return fmt.Errorf("could not check origins: %w", err) + } + case SwarmViewerCount: + log.Log(ctx, "got viewer count", "viewerCount", message) + if message.Server == swarm.NodeID { + // no infinite loops allowed + return nil + } + swarm.bus.SetViewerCount(message.Streamer, message.Server, message.Viewers) + log.Log(ctx, "set viewer count", "viewerCount", message) + return nil + default: + return fmt.Errorf("unknown message type: %s", reflect.TypeOf(rawMessage)) + } + return nil +} + +func decodeIrohMessage(key, value []byte) (any, error) { + keyStr := string(key) + if strings.HasPrefix(keyStr, "origin::") { + var originInfo SwarmOriginInfo + err := json.Unmarshal(value, &originInfo) + if err != nil { + return nil, fmt.Errorf("could not unmarshal origin info: %w", err) + } + return originInfo, nil + } + if strings.HasPrefix(keyStr, "viewers::") { + var viewerCount SwarmViewerCount + err := json.Unmarshal(value, &viewerCount) + if err != nil { + return nil, fmt.Errorf("could not unmarshal viewer count: %w", err) + } + return viewerCount, nil + } + return nil, fmt.Errorf("unknown key: %s", keyStr) +} + // subscribe to all streams func (swarm *IrohSwarm) startBusSubscribe(ctx context.Context) error { // start subscription first so we're buffering new origins @@ -218,6 +275,40 @@ func (swarm *IrohSwarm) startBusSubscribe(ctx context.Context) error { } } +func (swarm *IrohSwarm) startViewerCountSubscribe(ctx context.Context) error { + ch := swarm.bus.SubscribeToViewerCount() + for { + select { + case <-ctx.Done(): + return ctx.Err() + case msg := <-ch: + log.Log(ctx, "got viewer count update", "viewerCount", msg) + if msg.Origin != "local" { + continue + } + swarmMsg := SwarmViewerCount{ + Type: "place.stream.swarm.viewerCount", + Server: swarm.NodeID, + Streamer: msg.Streamer, + Viewers: msg.Count, + } + bs, err := json.Marshal(swarmMsg) + if err != nil { + log.Error(ctx, "could not marshal viewer count", "error", err) + continue + } + key := fmt.Sprintf("viewers::%s::%s", swarm.NodeID, msg.Streamer) + err = swarm.w.Put(nil, []byte(key), bs) + if err != nil { + log.Error(ctx, "could not put viewer count to swarm", "error", err) + continue + } + log.Log(ctx, "put viewer count to swarm", "viewerCount", msg) + + } + } +} + func (swarm *IrohSwarm) handleOriginMessage(ctx context.Context, view *streamplace.BroadcastDefs_BroadcastOriginView) error { origin, ok := view.Record.Val.(*streamplace.BroadcastOrigin) if !ok { @@ -247,6 +338,7 @@ func (swarm *IrohSwarm) handleOriginMessage(ctx context.Context, view *streampla } func (swarm *IrohSwarm) checkOrigins(ctx context.Context, streamer string, nodeID string) error { + ctx = log.WithLogValues(ctx, "streamer", streamer, "nodeID", nodeID, "func", "checkOrigins") err := swarm.cli.StreamIsAllowed(streamer) if err != nil { return fmt.Errorf("user %s is not allowlisted on this node: %w", streamer, err) @@ -290,9 +382,11 @@ func (swarm *IrohSwarm) checkOrigins(ctx context.Context, streamer string, nodeI log.Error(ctx, "could not subscribe to key", "error", err) return err } - swarm.activeSubs[streamer] = &OriginInfo{ - NodeID: nodeID, - Time: time.Now().Format(util.ISO8601), + swarm.activeSubs[streamer] = &SwarmOriginInfo{ + Type: "place.stream.swarm.originInfo", + NodeID: nodeID, + Time: time.Now().Format(util.ISO8601), + Streamer: streamer, } return nil } @@ -321,16 +415,18 @@ func (swarm *IrohSwarm) SendSegment(ctx context.Context, not *media.NewSegmentNo if !not.Local { return nil } - originInfo := OriginInfo{ - NodeID: swarm.NodeID, - Time: not.Segment.StartTime.Format(util.ISO8601), + originInfo := SwarmOriginInfo{ + Type: "place.stream.swarm.originInfo", + NodeID: swarm.NodeID, + Time: not.Segment.StartTime.Format(util.ISO8601), + Streamer: not.Segment.RepoDID, } bs, err := json.Marshal(originInfo) if err != nil { log.Error(ctx, "could not marshal origin info", "error", err) return err } - keyBs := []byte(not.Segment.RepoDID) + keyBs := []byte(fmt.Sprintf("origin::%s", not.Segment.RepoDID)) err = swarm.w.Put(nil, keyBs, bs) if err != nil { log.Error(ctx, "could not put segment to swarm", "error", err) diff --git a/pkg/spmetrics/spmetrics.go b/pkg/spmetrics/spmetrics.go index a7603a9e5..afbe28587 100644 --- a/pkg/spmetrics/spmetrics.go +++ b/pkg/spmetrics/spmetrics.go @@ -1,23 +1,16 @@ package spmetrics import ( - "context" "sync" - "time" "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/promauto" - "stream.place/streamplace/pkg/log" ) -const SessionExpireTime = 30 * time.Second //nolint:all - var viewersByStreamer = map[string]int{} var viewersByProtocol = map[string]int{} var viewersLock sync.RWMutex -var sessions = map[string]map[string]time.Time{} -var sessionsLock sync.RWMutex var Viewers = promauto.NewGaugeVec(prometheus.GaugeOpts{ Name: "streamplace_viewers", Help: "number of current viewers per user", @@ -110,39 +103,3 @@ func GetViewCount(user string) int { defer viewersLock.RUnlock() return viewersByStreamer[user] } - -func SessionSeen(user string, session string) { - now := time.Now() - go func() { - sessionsLock.Lock() - defer sessionsLock.Unlock() - if _, ok := sessions[user]; !ok { - sessions[user] = map[string]time.Time{} - } - if _, ok := sessions[user][session]; !ok { - log.Warn(context.TODO(), "ViewerInc", "user", user, "session", session) - ViewerInc(user, "hls") - } - sessions[user][session] = now - }() -} - -func ExpireSessions(ctx context.Context) error { - for { - select { - case <-ctx.Done(): - return nil - case <-time.After(5 * time.Second): - sessionsLock.Lock() - for user, sessions := range sessions { - for session, seen := range sessions { - if time.Since(seen) > SessionExpireTime { - delete(sessions, session) - ViewerDec(user, "hls") - } - } - } - sessionsLock.Unlock() - } - } -} diff --git a/pkg/statedb/xrpc_stream_event.go b/pkg/statedb/xrpc_stream_event.go index 1cb026043..73cfc7df0 100644 --- a/pkg/statedb/xrpc_stream_event.go +++ b/pkg/statedb/xrpc_stream_event.go @@ -30,7 +30,14 @@ func (ev *XrpcStreamEvent) ToCommitEvent() (*comatproto.SyncSubscribeRepos_Commi return commit, nil } +const CommitLockKey = "commit_lock" + func (state *StatefulDB) CreateCommitEvent(commit *comatproto.SyncSubscribeRepos_Commit, signedData string) error { + unlock, err := state.GetNamedLock(CommitLockKey) + if err != nil { + return err + } + defer unlock() prev, err := state.GetMostRecentCommitEvent(commit.Repo) if err != nil { return err