diff --git a/pkg/atproto/server_repo.go b/pkg/atproto/server_repo.go index 725c8392d..287de1945 100644 --- a/pkg/atproto/server_repo.go +++ b/pkg/atproto/server_repo.go @@ -6,6 +6,8 @@ import ( "fmt" "os" "path/filepath" + "sort" + "strings" "sync" "time" @@ -406,6 +408,40 @@ func ServerRepoMerkleProof(ctx context.Context, collection string, rkey string) return buf.Bytes(), nil } +// ServerRepoListCollections walks the server repo's MST and returns +// the distinct collection NSIDs currently holding at least one record. +// Used by com.atproto.repo.describeRepo to advertise what's actually +// in the repo rather than a hardcoded list. Returned collections are +// sorted lexicographically for stable output. +func ServerRepoListCollections(ctx context.Context) ([]string, error) { + serverRepoLock.Lock() + defer serverRepoLock.Unlock() + + r, _, err := OpenServerRepo(ctx) + if err != nil { + return nil, fmt.Errorf("ServerRepoListCollections: failed to open repo: %w", err) + } + seen := map[string]struct{}{} + err = r.ForEach(ctx, "", func(rpath string, _ cid.Cid) error { + // rpath is "/"; pull the prefix. + slash := strings.IndexByte(rpath, '/') + if slash <= 0 { + return nil + } + seen[rpath[:slash]] = struct{}{} + return nil + }) + if err != nil { + return nil, fmt.Errorf("ServerRepoListCollections: error iterating records: %w", err) + } + out := make([]string, 0, len(seen)) + for c := range seen { + out = append(out, c) + } + sort.Strings(out) + return out, nil +} + func ServerRepoListRecords(ctx context.Context, collection string, cursor string, limit int, repo string, reverse *bool) (*comatproto.RepoListRecords_Output, error) { serverRepoLock.Lock() defer serverRepoLock.Unlock() diff --git a/pkg/atproto/server_repo_test.go b/pkg/atproto/server_repo_test.go index 975f818c9..0b14a846c 100644 --- a/pkg/atproto/server_repo_test.go +++ b/pkg/atproto/server_repo_test.go @@ -66,6 +66,11 @@ func TestServerRepo(t *testing.T) { require.NoError(t, err) require.Len(t, listOut.Records, 1) + // ListCollections should report the collection we just wrote. + cols, err := ServerRepoListCollections(context.Background()) + require.NoError(t, err) + require.Equal(t, []string{constants.PLACE_STREAM_LIVE_VIEWERCOUNT}, cols) + // Merkle proof proof, err := ServerRepoMerkleProof(context.Background(), constants.PLACE_STREAM_LIVE_VIEWERCOUNT, "did:plc:abc123") require.NoError(t, err) @@ -93,5 +98,22 @@ func TestServerRepo(t *testing.T) { require.NoError(t, err) require.NotNil(t, out) + // After writing a record in a second collection, ListCollections + // should pick both up in sorted order. + origin := &streamplace.MediaOrigin{ + LexiconTypeID: constants.PLACE_STREAM_MEDIA_ORIGIN, + Blob: "babczxv...", + Size: 1234, + MimeType: "video/mp4", + } + err = CommitServerRepoRecord(context.Background(), &cli, constants.PLACE_STREAM_MEDIA_ORIGIN, "babczxv1", origin) + require.NoError(t, err) + cols, err = ServerRepoListCollections(context.Background()) + require.NoError(t, err) + require.Equal(t, []string{ + constants.PLACE_STREAM_LIVE_VIEWERCOUNT, + constants.PLACE_STREAM_MEDIA_ORIGIN, + }, cols) + handle.Close() } diff --git a/pkg/spxrpc/com_atproto_repo.go b/pkg/spxrpc/com_atproto_repo.go index 5969cca41..e61c330b4 100644 --- a/pkg/spxrpc/com_atproto_repo.go +++ b/pkg/spxrpc/com_atproto_repo.go @@ -17,7 +17,6 @@ import ( "go.opentelemetry.io/otel" "stream.place/streamplace/pkg/aqhttp" "stream.place/streamplace/pkg/atproto" - "stream.place/streamplace/pkg/constants" "stream.place/streamplace/pkg/log" ) @@ -92,13 +91,15 @@ func (s *Server) handleComAtprotoRepoDescribeRepo(ctx context.Context, repo stri } if s.isServerPDS(ctx) { + collections, err := atproto.ServerRepoListCollections(ctx) + if err != nil { + return nil, fmt.Errorf("list server repo collections: %w", err) + } return &comatproto.RepoDescribeRepo_Output{ - Handle: s.cli.ServerDID(), - Did: s.cli.ServerDID(), - DidDoc: atproto.DIDDoc(s.cli.ServerHost, atproto.ServerPubMultibase), - Collections: []string{ - constants.PLACE_STREAM_LIVE_VIEWERCOUNT, - }, + Handle: s.cli.ServerDID(), + Did: s.cli.ServerDID(), + DidDoc: atproto.DIDDoc(s.cli.ServerHost, atproto.ServerPubMultibase), + Collections: collections, HandleIsCorrect: true, }, nil }