From 478a7582631f85ccddb52308d7a645cae53cefc1 Mon Sep 17 00:00:00 2001 From: noelorin Date: Sun, 23 Aug 2026 13:59:20 +0800 Subject: [PATCH] feat(trunk): SRS as single trunk, trunk+distributors in config/events --- internal/server/gateway.go | 46 ++++++++--- internal/server/rooms.go | 77 +++++++++++++++++-- internal/server/service.go | 154 +++++++++++++++++++++---------------- 3 files changed, 193 insertions(+), 84 deletions(-) diff --git a/internal/server/gateway.go b/internal/server/gateway.go index 75632b3..d9548ed 100644 --- a/internal/server/gateway.go +++ b/internal/server/gateway.go @@ -81,12 +81,16 @@ func (s *Server) handleRoomEvents(w http.ResponseWriter, r *http.Request) { ch, unsub := s.hub.subscribe(room) defer unsub() - sendRoom := func(targets []*gen.StreamTarget) { - data, _ := protojson.Marshal(&gen.RoomEvent{Room: room, Targets: targets}) + sendRoom := func(targets []*gen.StreamTarget, dists []*gen.DistributionTarget) { + data, _ := protojson.Marshal(&gen.RoomEvent{Room: room, Targets: targets, Distributions: dists}) fmt.Fprintf(w, "event: room\ndata: %s\n\n", data) flusher.Flush() } - sendRoom(s.hub.targets(room)) + sendTyped := func(eventType string, raw json.RawMessage) { + fmt.Fprintf(w, "event: %s\ndata: %s\n\n", eventType, raw) + flusher.Flush() + } + sendRoom(s.hub.targets(room), s.hub.distributions(room)) ticker := time.NewTicker(20 * time.Second) defer ticker.Stop() @@ -95,9 +99,21 @@ func (s *Server) handleRoomEvents(w http.ResponseWriter, r *http.Request) { case <-r.Context().Done(): return case msg := <-ch: - sendRoom(decodeTargets(msg)) + var probe struct{ Type string `json:"type"` } + if err := json.Unmarshal(msg, &probe); err == nil && probe.Type != "" { + sendTyped(probe.Type, json.RawMessage(msg)) + } else { + targets, dists := decodeRoomEvent(msg) + if targets == nil { + targets = s.hub.targets(room) + } + if dists == nil { + dists = s.hub.distributions(room) + } + sendRoom(targets, dists) + } case <-ticker.C: - sendRoom(s.hub.targets(room)) + sendRoom(s.hub.targets(room), s.hub.distributions(room)) } } } @@ -147,17 +163,27 @@ func readBody(r *http.Request) []byte { } func decodeTargets(msg []byte) []*gen.StreamTarget { + targets, _ := decodeRoomEvent(msg) + return targets +} + +func decodeRoomEvent(msg []byte) ([]*gen.StreamTarget, []*gen.DistributionTarget) { var m struct { - Targets map[string]*gen.StreamTarget `json:"targets"` + Targets map[string]*gen.StreamTarget `json:"targets"` + Distributions map[string]*gen.DistributionTarget `json:"distributions"` } if err := json.Unmarshal(msg, &m); err != nil { - return nil + return nil, nil } - out := make([]*gen.StreamTarget, 0, len(m.Targets)) + targets := make([]*gen.StreamTarget, 0, len(m.Targets)) for _, t := range m.Targets { - out = append(out, t) + targets = append(targets, t) } - return out + dists := make([]*gen.DistributionTarget, 0, len(m.Distributions)) + for _, d := range m.Distributions { + dists = append(dists, d) + } + return targets, dists } func (s *Server) handleRoomsQuery(w http.ResponseWriter, r *http.Request) { diff --git a/internal/server/rooms.go b/internal/server/rooms.go index 8871477..bd84ba6 100644 --- a/internal/server/rooms.go +++ b/internal/server/rooms.go @@ -22,8 +22,9 @@ type roomHub struct { } type roomEntry struct { - targets map[string]*gen.StreamTarget // key: 后端名(cloudflare / srs) - subs map[chan []byte]struct{} + targets map[string]*gen.StreamTarget // key: 后端名(cloudflare / srs) + distributions map[string]*gen.DistributionTarget // key: distribution kind (srs-hls/cf/cdn) + subs map[chan []byte]struct{} } func newRoomHub() *roomHub { @@ -41,7 +42,7 @@ func newRoomHubWithDB(database *sql.DB) *roomHub { log.Printf("[hub] load from embedded db failed: %v", err) } else { for room, backends := range all { - e := &roomEntry{targets: map[string]*gen.StreamTarget{}, subs: map[chan []byte]struct{}{}} + e := &roomEntry{targets: map[string]*gen.StreamTarget{}, distributions: map[string]*gen.DistributionTarget{}, subs: map[chan []byte]struct{}{}} for backend, t := range backends { e.targets[backend] = t } @@ -61,7 +62,7 @@ func (h *roomHub) DB() *sql.DB { return h.db } func (h *roomHub) get(name string) *roomEntry { e, ok := h.rooms[name] if !ok { - e = &roomEntry{targets: map[string]*gen.StreamTarget{}, subs: map[chan []byte]struct{}{}} + e = &roomEntry{targets: map[string]*gen.StreamTarget{}, distributions: map[string]*gen.DistributionTarget{}, subs: map[chan []byte]struct{}{}} h.rooms[name] = e } return e @@ -71,7 +72,7 @@ func (h *roomHub) get(name string) *roomEntry { func (h *roomHub) createRoom(name string) { h.mu.Lock() if _, ok := h.rooms[name]; !ok { - h.rooms[name] = &roomEntry{targets: map[string]*gen.StreamTarget{}, subs: map[chan []byte]struct{}{}} + h.rooms[name] = &roomEntry{targets: map[string]*gen.StreamTarget{}, distributions: map[string]*gen.DistributionTarget{}, subs: map[chan []byte]struct{}{}} } h.mu.Unlock() } @@ -146,6 +147,9 @@ func (h *roomHub) list() []*gen.Room { for _, t := range e.targets { room.Targets = append(room.Targets, t) } + for _, d := range e.distributions { + room.Distributions = append(room.Distributions, d) + } out = append(out, room) } return out @@ -168,8 +172,9 @@ func (h *roomHub) subscribe(room string) (chan []byte, func()) { func (h *roomHub) broadcast(room string, e *roomEntry) { payload, _ := json.Marshal(map[string]interface{}{ - "room": room, - "targets": e.targets, + "room": room, + "targets": e.targets, + "distributions": e.distributions, }) for ch := range e.subs { select { @@ -179,6 +184,64 @@ func (h *roomHub) broadcast(room string, e *roomEntry) { } } +// broadcastEvent 向房间订阅者广播一个带类型的通用事件(playlist / playback / presence 等)。 +// 复用与拓扑相同的订阅通道,前端在同一个 SSE 连接上按 event type 分发处理。 +func (h *roomHub) broadcastEvent(room, eventType string, payload interface{}) { + h.mu.RLock() + e, ok := h.rooms[room] + if !ok { + h.mu.RUnlock() + return + } + envelope, _ := json.Marshal(map[string]interface{}{ + "type": eventType, + "room": room, + "payload": payload, + }) + for ch := range e.subs { + select { + case ch <- envelope: + default: + } + } + h.mu.RUnlock() +} + +func (h *roomHub) distributions(room string) []*gen.DistributionTarget { + h.mu.RLock() + defer h.mu.RUnlock() + e, ok := h.rooms[room] + if !ok { + return nil + } + out := make([]*gen.DistributionTarget, 0, len(e.distributions)) + for _, d := range e.distributions { + out = append(out, d) + } + return out +} + +func (h *roomHub) setDistribution(room string, d *gen.DistributionTarget) { + h.mu.Lock() + e := h.get(room) + if e.distributions == nil { + e.distributions = map[string]*gen.DistributionTarget{} + } + key := d.GetKind().String() + e.distributions[key] = d + h.broadcast(room, e) + h.mu.Unlock() +} + +func (h *roomHub) clearDistributions(room string) { + h.mu.Lock() + if e, ok := h.rooms[room]; ok { + e.distributions = map[string]*gen.DistributionTarget{} + h.broadcast(room, e) + } + h.mu.Unlock() +} + func (h *roomHub) removeRoom(name string) { h.mu.Lock() defer h.mu.Unlock() diff --git a/internal/server/service.go b/internal/server/service.go index acb52d9..8081aeb 100644 --- a/internal/server/service.go +++ b/internal/server/service.go @@ -19,8 +19,16 @@ type Service struct { srsP *srs.Provider hub *roomHub auth *auth.Manager + distributionManager DistributionManagerIface } +// DistributionManagerIface abstracts distribution fanout for testability +type DistributionManagerIface interface { + Ensure(ctx context.Context, room, stream string) error +} + +func (s *Service) SetDistributionManager(m DistributionManagerIface) { s.distributionManager = m } + func NewService(cfg *config.Config, cf *cloudflare.Provider, srsP *srs.Provider, hub *roomHub) *Service { return &Service{cfg: cfg, cf: cf, srsP: srsP, hub: hub} } @@ -64,6 +72,26 @@ func (s *Service) GetConfig(ctx context.Context, req *gen.GetConfigRequest) (*ge resp.Backends = append(resp.Backends, s.srsP.BackendInfo()) } } + // SRS is the trunk + resp.Trunk = s.srsP.BackendInfo() + resp.Trunk.Primary = true + // Build distributors list (capability matrix) + for _, d := range s.cfg.DistributorList() { + switch d { + case "srs-hls": + resp.Distributors = append(resp.Distributors, &gen.DistributorInfo{ + Kind: gen.DistributionKind_DISTRIBUTION_KIND_SRS_HLS, Name: "SRS HLS", Configured: true, Enabled: true, + }) + case "cf": + resp.Distributors = append(resp.Distributors, &gen.DistributorInfo{ + Kind: gen.DistributionKind_DISTRIBUTION_KIND_CF_SFU, Name: "Cloudflare SFU", Configured: s.cf.Configured(), Enabled: s.cf.Configured(), + }) + case "cdn": + resp.Distributors = append(resp.Distributors, &gen.DistributorInfo{ + Kind: gen.DistributionKind_DISTRIBUTION_KIND_CDN, Name: s.cfg.CDNName, Configured: s.cfg.CDNEnabled(), Enabled: s.cfg.CDNEnabled(), + }) + } + } return resp, nil } @@ -82,57 +110,30 @@ func (s *Service) Publish(ctx context.Context, req *gen.PublishRequest) (*gen.Pu if room == "" { return nil, fmt.Errorf("room required") } - backend := backendName(req.GetBackend()) - if backend == "" { - return nil, fmt.Errorf("unknown backend") - } identity := req.GetIdentity() if identity == "" { identity = randomID() } - switch backend { - case "cloudflare": - if !s.cf.Configured() { - return nil, fmt.Errorf("cloudflare realtime not configured (set CF_APP_ID / CF_APP_SECRET)") - } - sessionID, err := s.cf.Client().CreateSession(room) - if err != nil { - return nil, fmt.Errorf("cf create session: %w", err) - } - info, err := s.cf.Client().GetSession(sessionID) - if err != nil { - return nil, fmt.Errorf("cf get session: %w", err) - } - target := &gen.StreamTarget{ - Backend: gen.BackendKind_BACKEND_KIND_CLOUDFLARE, - SessionId: sessionID, - PublishedAt: time.Now().Unix(), - } - s.hub.setTarget(room, "cloudflare", target) - return &gen.PublishResponse{ - SessionId: sessionID, - IceServers: iceServersToProto(info.IceServers), - Target: target, - }, nil - case "srs": - stream := "live-" + room - token, _ := signStreamToken(s.cfg.TokenSecret, stream, identity, "publish", 2*time.Hour) - target := &gen.StreamTarget{ - Backend: gen.BackendKind_BACKEND_KIND_SRS, - Stream: stream, - PublishToken: token, - Url: fmt.Sprintf("/rtc/v1/whep/?app=%s&stream=%s", s.srsP.App(), stream), - PublishedAt: time.Now().Unix(), - } - s.hub.setTarget(room, "srs", target) - return &gen.PublishResponse{ - Stream: stream, - PublishToken: token, - IceServers: []*gen.IceServer{{Urls: []string{s.cf.Stun()}}}, - Target: target, - }, nil + // SRS is the only trunk: always publish to SRS regardless of requested backend + stream := "live-" + room + token, _ := signStreamToken(s.cfg.TokenSecret, stream, identity, "publish", 2*time.Hour) + target := &gen.StreamTarget{ + Backend: gen.BackendKind_BACKEND_KIND_SRS, + Stream: stream, + PublishToken: token, + Url: fmt.Sprintf("/rtc/v1/whep/?app=%s&stream=%s", s.srsP.App(), stream), + PublishedAt: time.Now().Unix(), } - return nil, fmt.Errorf("unsupported backend") + s.hub.setTarget(room, "srs", target) + if s.distributionManager != nil { + _ = s.distributionManager.Ensure(ctx, room, stream) + } + return &gen.PublishResponse{ + Stream: stream, + PublishToken: token, + IceServers: []*gen.IceServer{{Urls: []string{s.cf.Stun()}}}, + Target: target, + }, nil } func (s *Service) Subscribe(ctx context.Context, req *gen.SubscribeRequest) (*gen.SubscribeResponse, error) { @@ -143,16 +144,20 @@ func (s *Service) Subscribe(ctx context.Context, req *gen.SubscribeRequest) (*ge if room == "" { return nil, fmt.Errorf("room required") } - backend := backendName(req.GetBackend()) - if backend == "" { - return nil, fmt.Errorf("unknown backend") - } - pub := s.findTarget(room, req.GetBackend()) + // Always read from SRS trunk; backend param is kept for compatibility + pub := s.findTarget(room, gen.BackendKind_BACKEND_KIND_SRS) if pub == nil { - return nil, fmt.Errorf("room %q not live on %s", room, backend) + return nil, fmt.Errorf("room %q not live", room) } - switch backend { - case "cloudflare": + // If caller explicitly requested Cloudflare backend, still need viewer session via CF + // Check requested backend to decide response shape + reqBackend := req.GetBackend() + if reqBackend == gen.BackendKind_BACKEND_KIND_CLOUDFLARE { + if !s.cf.Configured() { + return nil, fmt.Errorf("cloudflare realtime not configured") + } + // For CF distribution, viewer subscribes to publisher session + // Publisher session is not trunk session; use trunk stream as identity viewerSession, err := s.cf.Client().CreateSession(room) if err != nil { return nil, fmt.Errorf("cf create viewer session: %w", err) @@ -163,16 +168,14 @@ func (s *Service) Subscribe(ctx context.Context, req *gen.SubscribeRequest) (*ge } return &gen.SubscribeResponse{ SessionId: viewerSession, - PublisherSessionId: pub.GetSessionId(), + PublisherSessionId: pub.GetStream(), // trunk stream as publisher id for CF relay IceServers: iceServersToProto(info.IceServers), }, nil - case "srs": - return &gen.SubscribeResponse{ - Stream: pub.GetStream(), - IceServers: []*gen.IceServer{{Urls: []string{s.cf.Stun()}}}, - }, nil } - return nil, fmt.Errorf("unsupported backend") + return &gen.SubscribeResponse{ + Stream: pub.GetStream(), + IceServers: []*gen.IceServer{{Urls: []string{s.cf.Stun()}}}, + }, nil } func (s *Service) StopStream(ctx context.Context, req *gen.StopStreamRequest) (*gen.StopStreamResponse, error) { @@ -180,10 +183,14 @@ func (s *Service) StopStream(ctx context.Context, req *gen.StopStreamRequest) (* return nil, err } room := req.GetRoom() - backend := backendName(req.GetBackend()) - if room == "" || backend == "" { + if room == "" { return nil, fmt.Errorf("room and backend required") } + // SRS trunk is the only target to remove; also clean CF sessions if any + backend := backendName(req.GetBackend()) + if backend == "" { + backend = "srs" + } if backend == "cloudflare" { for _, t := range s.hub.targets(room) { if t.GetBackend() == gen.BackendKind_BACKEND_KIND_CLOUDFLARE && t.GetSessionId() != "" { @@ -191,7 +198,13 @@ func (s *Service) StopStream(ctx context.Context, req *gen.StopStreamRequest) (* } } } - s.hub.removeTarget(room, backend) + // Always remove SRS trunk if present + s.hub.removeTarget(room, "srs") + if backend != "srs" { + s.hub.removeTarget(room, backend) + } + // Also clear distributions + s.hub.clearDistributions(room) return &gen.StopStreamResponse{Ok: true}, nil } @@ -202,7 +215,7 @@ func (s *Service) WatchRoom(req *gen.WatchRoomRequest, stream gen.SyncLive_Watch room := req.GetRoom() ch, unsub := s.hub.subscribe(room) defer unsub() - if err := stream.Send(&gen.RoomEvent{Room: room, Targets: s.hub.targets(room)}); err != nil { + if err := stream.Send(&gen.RoomEvent{Room: room, Targets: s.hub.targets(room), Distributions: s.hub.distributions(room)}); err != nil { return err } ticker := time.NewTicker(20 * time.Second) @@ -212,11 +225,18 @@ func (s *Service) WatchRoom(req *gen.WatchRoomRequest, stream gen.SyncLive_Watch case <-stream.Context().Done(): return nil case msg := <-ch: - if err := stream.Send(&gen.RoomEvent{Room: room, Targets: decodeTargets(msg)}); err != nil { + targets, dists := decodeRoomEvent(msg) + if targets == nil { + targets = s.hub.targets(room) + } + if dists == nil { + dists = s.hub.distributions(room) + } + if err := stream.Send(&gen.RoomEvent{Room: room, Targets: targets, Distributions: dists}); err != nil { return err } case <-ticker.C: - if err := stream.Send(&gen.RoomEvent{Room: room, Targets: s.hub.targets(room)}); err != nil { + if err := stream.Send(&gen.RoomEvent{Room: room, Targets: s.hub.targets(room), Distributions: s.hub.distributions(room)}); err != nil { return err } }