diff --git a/pkg/cmd/whep.go b/pkg/cmd/whep.go index 827973f7a..bfc40a112 100644 --- a/pkg/cmd/whep.go +++ b/pkg/cmd/whep.go @@ -47,6 +47,7 @@ type WHEPClient struct { Endpoint string Count int FreezeAfter time.Duration + Stats []map[string]*TrackStats } type WHEPConnection struct { @@ -57,7 +58,14 @@ type WHEPConnection struct { Done func() <-chan struct{} } -func (w *WHEPClient) StartWHEPConnection(ctx context.Context) (*WHEPConnection, error) { +type TrackStats struct { + Total int + lastTotal int + lastUpdate time.Time + mu sync.Mutex +} + +func (w *WHEPClient) StartWHEPConnection(ctx context.Context, stats map[string]*TrackStats) (*WHEPConnection, error) { // Prepare the configuration config := webrtc.Configuration{} @@ -68,19 +76,6 @@ func (w *WHEPClient) StartWHEPConnection(ctx context.Context) (*WHEPConnection, return nil, err } - // Track statistics - type trackStats struct { - total int - lastTotal int - lastUpdate time.Time - mu sync.Mutex - } - - stats := map[string]*trackStats{ - "video": {lastUpdate: time.Now()}, - "audio": {lastUpdate: time.Now()}, - } - // Create a ticker to print combined bitrate every 5 seconds ticker := time.NewTicker(5 * time.Second) @@ -102,8 +97,8 @@ func (w *WHEPClient) StartWHEPConnection(ctx context.Context) (*WHEPConnection, videoElapsed := currentTime.Sub(videoStats.lastUpdate).Seconds() audioElapsed := currentTime.Sub(audioStats.lastUpdate).Seconds() - videoBytes := videoStats.total - videoStats.lastTotal - audioBytes := audioStats.total - audioStats.lastTotal + videoBytes := videoStats.Total - videoStats.lastTotal + audioBytes := audioStats.Total - audioStats.lastTotal videoBitrate := float64(videoBytes) * 8 / videoElapsed / 1000 // kbps audioBitrate := float64(audioBytes) * 8 / audioElapsed / 1000 // kbps @@ -114,9 +109,9 @@ func (w *WHEPClient) StartWHEPConnection(ctx context.Context) (*WHEPConnection, "total", fmt.Sprintf("%.2f kbps", videoBitrate+audioBitrate)) // Update last values - videoStats.lastTotal = videoStats.total + videoStats.lastTotal = videoStats.Total videoStats.lastUpdate = currentTime - audioStats.lastTotal = audioStats.total + audioStats.lastTotal = audioStats.Total audioStats.lastUpdate = currentTime // Unlock stats @@ -156,7 +151,7 @@ func (w *WHEPClient) StartWHEPConnection(ctx context.Context) (*WHEPConnection, } trackStat.mu.Lock() - trackStat.total += len(rtp.Payload) + trackStat.Total += len(rtp.Payload) trackStat.mu.Unlock() } }) @@ -252,14 +247,20 @@ func (w *WHEPClient) StartWHEPConnection(ctx context.Context) (*WHEPConnection, } func (w *WHEPClient) WHEP(ctx context.Context) error { + w.Stats = []map[string]*TrackStats{} ctx, cancel := context.WithCancel(ctx) defer cancel() conns := make([]*WHEPConnection, w.Count) g := &errgroup.Group{} for i := 0; i < w.Count; i++ { + stats := map[string]*TrackStats{ + "video": {lastUpdate: time.Now()}, + "audio": {lastUpdate: time.Now()}, + } + w.Stats = append(w.Stats, stats) g.Go(func() error { - conn, err := w.StartWHEPConnection(ctx) + conn, err := w.StartWHEPConnection(ctx, stats) if err != nil { return err } diff --git a/pkg/config/config.go b/pkg/config/config.go index 5f37553b2..1bf300507 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -129,6 +129,7 @@ type CLI struct { IrohTopic string DID string DisableIrohRelay bool + DevAccountCreds map[string]string } // ContentFilters represents the content filtering configuration @@ -214,6 +215,7 @@ func (cli *CLI) NewFlagSet(name string) *flag.FlagSet { cli.StringSliceFlag(fs, &cli.Tickets, "tickets", "[]", "tickets to join the swarm with") fs.StringVar(&cli.IrohTopic, "iroh-topic", "", "topic to use for the iroh swarm (must be 32 bytes in hex)") fs.BoolVar(&cli.DisableIrohRelay, "disable-iroh-relay", false, "disable the iroh relay") + cli.KVSliceFlag(fs, &cli.DevAccountCreds, "dev-account-creds", "", "(FOR DEVELOPMENT ONLY) did=password pairs for logging into test accounts without oauth") lpFlags := flag.NewFlagSet("livepeer", flag.ContinueOnError) _ = starter.NewLivepeerConfig(lpFlags) @@ -542,6 +544,25 @@ func (cli *CLI) StringSliceFlag(fs *flag.FlagSet, dest *[]string, name, defaultV }) } +func (cli *CLI) KVSliceFlag(fs *flag.FlagSet, dest *map[string]string, name, defaultValue, usage string) { + *dest = map[string]string{} + usage = fmt.Sprintf(`%s (default: "%s")`, usage, *dest) + fs.Func(name, usage, func(s string) error { + if s == "" { + return nil + } + pairs := strings.Split(s, ",") + for _, pair := range pairs { + parts := strings.Split(pair, "=") + if len(parts) != 2 { + return fmt.Errorf("invalid kv flag: %s", pair) + } + (*dest)[parts[0]] = parts[1] + } + return nil + }) +} + func (cli *CLI) JSONFlag(fs *flag.FlagSet, dest any, name, defaultValue, usage string) { usage = fmt.Sprintf(`%s (default: "%s")`, usage, defaultValue) fs.Func(name, usage, func(s string) error { diff --git a/pkg/crypto/spkey/spkey.go b/pkg/crypto/spkey/spkey.go index 91c0c0ce1..f23ae3116 100644 --- a/pkg/crypto/spkey/spkey.go +++ b/pkg/crypto/spkey/spkey.go @@ -30,7 +30,7 @@ func GenerateStreamKeyForDID(did string) (string, *atcrypto.PublicKeyK256, error } didBytes := []byte(did) combinedBytes := append(priv.Bytes(), didBytes...) - multibaseKey := base58.Encode(combinedBytes) + multibaseKey := "z" + base58.Encode(combinedBytes) return multibaseKey, pub, nil } diff --git a/pkg/devenv/devenv.go b/pkg/devenv/devenv.go index a184df6da..912536392 100644 --- a/pkg/devenv/devenv.go +++ b/pkg/devenv/devenv.go @@ -25,8 +25,9 @@ import ( ) type DevEnv struct { - PDSURL string `json:"pds-url"` - PLCURL string `json:"plc-url"` + PDSURL string `json:"pds-url"` + PLCURL string `json:"plc-url"` + Accounts []*DevEnvAccount `json:"accounts"` } func WithDevEnv(t *testing.T) *DevEnv { @@ -55,6 +56,7 @@ func WithDevEnv(t *testing.T) *DevEnv { t.Logf("Error unmarshalling dev-env stdout: %v", err) t.FailNow() } + env.Accounts = []*DevEnvAccount{} go func() { scanner := bufio.NewScanner(stdout) @@ -124,14 +126,15 @@ func (d *DevEnv) CreateAccount(t *testing.T) *DevEnvAccount { Handle: out.Handle, }, } - - return &DevEnvAccount{ + acct := &DevEnvAccount{ Handle: out.Handle, Email: email, Password: password, DID: out.Did, XRPC: xrpcc, } + d.Accounts = append(d.Accounts, acct) + return acct } // Custom RoundTripper for intercepting .test domain requests diff --git a/pkg/director/stream_session.go b/pkg/director/stream_session.go index 0cd485dee..52ada5bf5 100644 --- a/pkg/director/stream_session.go +++ b/pkg/director/stream_session.go @@ -7,13 +7,14 @@ import ( "sync" "time" - "github.com/bluesky-social/indigo/api/atproto" + comatproto "github.com/bluesky-social/indigo/api/atproto" "github.com/bluesky-social/indigo/api/bsky" lexutil "github.com/bluesky-social/indigo/lex/util" "github.com/bluesky-social/indigo/util" "github.com/bluesky-social/indigo/xrpc" "github.com/streamplace/oatproxy/pkg/oatproxy" "golang.org/x/sync/errgroup" + "stream.place/streamplace/pkg/aqhttp" "stream.place/streamplace/pkg/aqtime" "stream.place/streamplace/pkg/bus" "stream.place/streamplace/pkg/config" @@ -370,7 +371,7 @@ func (ss *StreamSession) UpdateStatus(ctx context.Context, repoDID string) error } var swapRecord *string - getOutput := atproto.RepoGetRecord_Output{} + getOutput := comatproto.RepoGetRecord_Output{} err = client.Do(ctx, xrpc.Query, "application/json", "com.atproto.repo.getRecord", map[string]any{ "repo": repoDID, "collection": "app.bsky.actor.status", @@ -390,14 +391,14 @@ func (ss *StreamSession) UpdateStatus(ctx context.Context, repoDID string) error swapRecord = getOutput.Cid } - inp := atproto.RepoPutRecord_Input{ + inp := comatproto.RepoPutRecord_Input{ Collection: "app.bsky.actor.status", Record: &lexutil.LexiconTypeDecoder{Val: &status}, Rkey: "self", Repo: repoDID, SwapRecord: swapRecord, } - out := atproto.RepoPutRecord_Output{} + out := comatproto.RepoPutRecord_Output{} ss.lastStatusCID = &out.Cid @@ -421,13 +422,13 @@ func (ss *StreamSession) DeleteStatus(repoDID string) error { log.Debug(ctx, "no status cid to delete") return nil } - inp := atproto.RepoDeleteRecord_Input{ + inp := comatproto.RepoDeleteRecord_Input{ Collection: "app.bsky.actor.status", Rkey: "self", Repo: repoDID, } inp.SwapRecord = ss.lastStatusCID - out := atproto.RepoDeleteRecord_Output{} + out := comatproto.RepoDeleteRecord_Output{} client, err := ss.GetClientByDID(repoDID) if err != nil { @@ -470,7 +471,7 @@ func (ss *StreamSession) UpdateBroadcastOrigin(ctx context.Context) error { rkey := fmt.Sprintf("%s::did:web:%s", ss.repoDID, ss.cli.ServerHost) var swapRecord *string - getOutput := atproto.RepoGetRecord_Output{} + getOutput := comatproto.RepoGetRecord_Output{} err = client.Do(ctx, xrpc.Query, "application/json", "com.atproto.repo.getRecord", map[string]any{ "repo": ss.repoDID, "collection": "place.stream.broadcast.origin", @@ -490,14 +491,14 @@ func (ss *StreamSession) UpdateBroadcastOrigin(ctx context.Context) error { swapRecord = getOutput.Cid } - inp := atproto.RepoPutRecord_Input{ + inp := comatproto.RepoPutRecord_Input{ Collection: "place.stream.broadcast.origin", Record: &lexutil.LexiconTypeDecoder{Val: &origin}, Rkey: rkey, Repo: ss.repoDID, SwapRecord: swapRecord, } - out := atproto.RepoPutRecord_Output{} + out := comatproto.RepoPutRecord_Output{} err = client.Do(ctx, xrpc.Procedure, "application/json", "com.atproto.repo.putRecord", map[string]any{}, inp, &out) if err != nil { @@ -621,6 +622,40 @@ type XRPCClient interface { } func (ss *StreamSession) GetClientByDID(did string) (XRPCClient, error) { + password, ok := ss.cli.DevAccountCreds[did] + if ok { + repo, err := ss.mod.GetRepoByHandleOrDID(did) + if err != nil { + return nil, fmt.Errorf("could not get repo by did: %w", err) + } + if repo == nil { + return nil, fmt.Errorf("repo not found for did: %s", did) + } + anonXRPCC := &xrpc.Client{ + Host: repo.PDS, + Client: &aqhttp.Client, + } + session, err := comatproto.ServerCreateSession(context.Background(), anonXRPCC, &comatproto.ServerCreateSession_Input{ + Identifier: repo.DID, + Password: password, + }) + if err != nil { + return nil, fmt.Errorf("could not create session: %w", err) + } + + log.Warn(context.Background(), "created session for dev account", "did", repo.DID, "handle", repo.Handle, "pds", repo.PDS) + + return &xrpc.Client{ + Host: repo.PDS, + Client: &aqhttp.Client, + Auth: &xrpc.AuthInfo{ + Did: repo.DID, + AccessJwt: session.AccessJwt, + RefreshJwt: session.RefreshJwt, + Handle: repo.Handle, + }, + }, nil + } session, err := ss.statefulDB.GetSessionByDID(ss.repoDID) if err != nil { return nil, fmt.Errorf("could not get OAuth session for repoDID: %w", err) diff --git a/pkg/multitest/multitest_test.go b/pkg/multitest/multitest_test.go index 11d34e901..ae54c4353 100644 --- a/pkg/multitest/multitest_test.go +++ b/pkg/multitest/multitest_test.go @@ -3,6 +3,7 @@ package multitest import ( "context" "fmt" + "net/http" "os" "os/exec" "path/filepath" @@ -15,18 +16,22 @@ import ( lexutil "github.com/bluesky-social/indigo/lex/util" "github.com/bluesky-social/indigo/util" "github.com/stretchr/testify/require" + "golang.org/x/sync/errgroup" + "stream.place/streamplace/pkg/cmd" "stream.place/streamplace/pkg/crypto/spkey" "stream.place/streamplace/pkg/devenv" + "stream.place/streamplace/pkg/gstinit" "stream.place/streamplace/pkg/log" "stream.place/streamplace/pkg/streamplace" ) func TestMultinodeSyndication(t *testing.T) { + gstinit.InitGST() dev := devenv.WithDevEnv(t) - startStreamplaceNode(t, dev) - // startStreamplaceNode(t, dev) acct := dev.CreateAccount(t) - _, pub, err := spkey.GenerateStreamKeyForDID(acct.DID) + node1 := startStreamplaceNode(t, dev) + node2 := startStreamplaceNode(t, dev) + priv, pub, err := spkey.GenerateStreamKeyForDID(acct.DID) require.NoError(t, err) createdBy := "multitest" streamKey := streamplace.Key{ @@ -41,6 +46,37 @@ func TestMultinodeSyndication(t *testing.T) { }) require.NoError(t, err) log.Log(context.Background(), "created stream key", "did", acct.DID, "pub", pub.DIDKey()) + whip := &cmd.WHIPClient{ + StreamKey: priv, + File: "/home/iameli/testvids/RocketLeague_1h55m_1sGOP_1080p60_NoBframes.mp4", + Endpoint: fmt.Sprintf("http://%s", node1.Env["SP_HTTP_ADDR"]), + Count: 1, + } + + whep := &cmd.WHEPClient{ + Endpoint: fmt.Sprintf("http://%s/api/playback/%s/webrtc", node2.Env["SP_HTTP_ADDR"], acct.DID), + Count: 1, + } + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + g, ctx := errgroup.WithContext(ctx) + g.Go(func() error { + return whip.WHIP(ctx) + }) + g.Go(func() error { + return whep.WHEP(ctx) + }) + + <-ctx.Done() + + err = g.Wait() + require.ErrorIs(t, err, context.DeadlineExceeded) + stats := whep.Stats[0] + videoStats := stats["video"] + audioStats := stats["audio"] + require.Greater(t, videoStats.Total, 0) + require.Greater(t, audioStats.Total, 0) } var currentPort = 10000 @@ -50,14 +86,23 @@ func nextPort() int { return currentPort } -func startStreamplaceNode(t *testing.T, dev *devenv.DevEnv) { +type TestNode struct { + Env map[string]string +} + +func startStreamplaceNode(t *testing.T, dev *devenv.DevEnv) *TestNode { dataDir := t.TempDir() + devAccountCreds := []string{} + for _, acct := range dev.Accounts { + devAccountCreds = append(devAccountCreds, fmt.Sprintf("%s=%s", acct.DID, acct.Password)) + } env := map[string]string{ - "SP_HTTP_ADDR": fmt.Sprintf(":%d", nextPort()), - "SP_HTTP_INTERNAL_ADDR": fmt.Sprintf(":%d", nextPort()), + "SP_HTTP_ADDR": fmt.Sprintf("127.0.0.1:%d", nextPort()), + "SP_HTTP_INTERNAL_ADDR": fmt.Sprintf("127.0.0.1:%d", nextPort()), "SP_RELAY_HOST": strings.ReplaceAll(dev.PDSURL, "http://", "ws://"), "SP_PLC_URL": dev.PLCURL, "SP_DATA_DIR": dataDir, + "SP_DEV_ACCOUNT_CREDS": strings.Join(devAccountCreds, ","), } _, file, _, _ := runtime.Caller(0) abs, err := filepath.Abs(filepath.Join(filepath.Dir(file), "..", "..", "build-linux-amd64", "streamplace")) @@ -78,4 +123,20 @@ func startStreamplaceNode(t *testing.T, dev *devenv.DevEnv) { _, err = cmd.Process.Wait() require.NoError(t, err) }) + // Wait for the streamplace node to be ready by polling the health endpoint + healthz := fmt.Sprintf("http://%s/api/healthz", env["SP_HTTP_ADDR"]) + client := &http.Client{Timeout: 2 * time.Second} + for { + resp, err := client.Get(healthz) + if err == nil { + defer resp.Body.Close() + if resp.StatusCode == 200 { + break + } + } + time.Sleep(200 * time.Millisecond) + } + return &TestNode{ + Env: env, + } }