diff --git a/pkg/atproto/lexicon_repo_queries.go b/pkg/atproto/lexicon_repo_queries.go new file mode 100644 index 000000000..a3d09a00b --- /dev/null +++ b/pkg/atproto/lexicon_repo_queries.go @@ -0,0 +1,128 @@ +package atproto + +import ( + "bytes" + "context" + "fmt" + "sync" + + "github.com/bluesky-social/indigo/carstore" + lexutil "github.com/bluesky-social/indigo/lex/util" + "github.com/bluesky-social/indigo/repo" + "github.com/bluesky-social/indigo/util" + "github.com/ipfs/go-cid" + cbor "github.com/ipfs/go-ipld-cbor" + "github.com/ipld/go-car" + "stream.place/streamplace/pkg/log" + + comatproto "github.com/bluesky-social/indigo/api/atproto" +) + +var repoLock sync.Mutex + +func LexiconRepoMerkleProof(ctx context.Context, collection string, rkey string) ([]byte, error) { + repoLock.Lock() + defer repoLock.Unlock() + + _, robs, err := OpenLexiconRepo(ctx) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to open repo: %w", err) + } + + bs := util.NewLoggingBstore(robs) + + root, err := CarStore.GetUserRepoHead(ctx, RepoUser) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to get user repo head: %w", err) + } + + log.Warn(ctx, "got root", "root", root.String()) + + r, err := repo.OpenRepo(ctx, bs, root) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to open repo: %w", err) + } + + _, _, err = r.GetRecordBytes(ctx, collection+"/"+rkey) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to get record bytes: %w", err) + } + + blocks := bs.GetLoggedBlocks() + + buf := new(bytes.Buffer) + hb, err := cbor.DumpObject(&car.CarHeader{ + Roots: []cid.Cid{root}, + Version: 1, + }) + if err != nil { + return nil, fmt.Errorf("failed to dump car header: %w", err) + } + if _, err := carstore.LdWrite(buf, hb); err != nil { + return nil, err + } + + for _, blk := range blocks { + log.Warn(ctx, "writing block", "cid", blk.Cid().String(), "version", blk.Cid().Version()) + if _, err := carstore.LdWrite(buf, blk.Cid().Bytes(), blk.RawData()); err != nil { + return nil, err + } + } + + return buf.Bytes(), nil +} + +func LexiconRepoListRecords(ctx context.Context, collection string, cursor string, limit int, repo string, reverse *bool) (*comatproto.RepoListRecords_Output, error) { + repoLock.Lock() + defer repoLock.Unlock() + + r, ses, err := OpenLexiconRepo(ctx) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoListRecords: failed to open repo: %w", err) + } + out := &comatproto.RepoListRecords_Output{ + Records: []*comatproto.RepoListRecords_Record{}, + } + err = r.ForEach(ctx, "", func(rkey string, c cid.Cid) error { + val, err := GetRecordCBOR(ctx, ses, c, collection, rkey) + if err != nil { + return fmt.Errorf("handleComAtprotoRepoListRecords: failed to get record for collection %q, rkey %q: %w", collection, rkey, err) + } + log.Warn(ctx, "got record", "rkey", rkey, "cid", c.String()) + out.Records = append(out.Records, &comatproto.RepoListRecords_Record{ + Uri: fmt.Sprintf("at://%s/%s", repo, rkey), + Cid: c.String(), + Value: &lexutil.LexiconTypeDecoder{Val: val}, + }) + + return nil + }) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoListRecords: error iterating records for collection %q: %w", collection, err) + } + return out, nil +} + +func LexiconRepoGetRecord(ctx context.Context, repo string, collection string, rkey string) (*comatproto.RepoGetRecord_Output, error) { + repoLock.Lock() + defer repoLock.Unlock() + + r, ses, err := OpenLexiconRepo(ctx) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to open repo: %w", err) + } + outCID, _, err := r.GetRecord(ctx, fmt.Sprintf("%s/%s", collection, rkey)) + if err != nil { + return nil, err + } + rec, err := GetRecordCBOR(ctx, ses, outCID, collection, rkey) + if err != nil { + return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to get record: %w", err) + } + str := outCID.String() + return &comatproto.RepoGetRecord_Output{ + Uri: fmt.Sprintf("at://%s/%s/%s", repo, collection, rkey), + Cid: &str, + Value: &lexutil.LexiconTypeDecoder{Val: rec}, + }, nil +} diff --git a/pkg/atproto/lexicon_repo_queries_test.go b/pkg/atproto/lexicon_repo_queries_test.go new file mode 100644 index 000000000..100b4c0e1 --- /dev/null +++ b/pkg/atproto/lexicon_repo_queries_test.go @@ -0,0 +1,53 @@ +package atproto + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + "golang.org/x/sync/errgroup" + "stream.place/streamplace/pkg/config" + "stream.place/streamplace/pkg/model" + "stream.place/streamplace/pkg/statedb" +) + +func TestLexiconRepoConcurrentAccess(t *testing.T) { + cli := config.CLI{ + BroadcasterHost: "example.com", + DBURL: ":memory:", + DataDir: t.TempDir(), + } + mod, err := model.MakeDB(":memory:") + require.NoError(t, err) + state, err := statedb.MakeDB(context.Background(), &cli, nil, mod) + require.NoError(t, err) + handle, err := MakeLexiconRepo(context.Background(), &cli, mod, state) + require.NoError(t, err) + handle.Close() + ctx := context.Background() + collection := "com.atproto.lexicon.schema" + + g, ctx := errgroup.WithContext(ctx) + for i := 0; i < 10; i++ { + g.Go(func() error { + res, err := LexiconRepoMerkleProof(ctx, collection, "place.stream.chat.message") + require.NoError(t, err) + require.NotNil(t, res) + return nil + }) + g.Go(func() error { + res, err := LexiconRepoListRecords(ctx, collection, "", 10, "did:web:example.com", nil) + require.NoError(t, err) + require.NotNil(t, res) + return nil + }) + g.Go(func() error { + res, err := LexiconRepoGetRecord(ctx, "did:web:example.com", collection, "place.stream.chat.message") + require.NoError(t, err) + require.NotNil(t, res) + return nil + }) + } + err = g.Wait() + require.NoError(t, err) +} diff --git a/pkg/atproto/lexicon_repo_test.go b/pkg/atproto/lexicon_repo_test.go index a70f6c094..f014701aa 100644 --- a/pkg/atproto/lexicon_repo_test.go +++ b/pkg/atproto/lexicon_repo_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/require" "stream.place/streamplace/lexicons" + "stream.place/streamplace/pkg/config" "stream.place/streamplace/pkg/model" "stream.place/streamplace/pkg/statedb" diff --git a/pkg/spxrpc/com_atproto_repo.go b/pkg/spxrpc/com_atproto_repo.go index 0d68905b8..f29ce935e 100644 --- a/pkg/spxrpc/com_atproto_repo.go +++ b/pkg/spxrpc/com_atproto_repo.go @@ -9,11 +9,9 @@ import ( "net/http" "strings" - comatprototypes "github.com/bluesky-social/indigo/api/atproto" - lexutil "github.com/bluesky-social/indigo/lex/util" + comatproto "github.com/bluesky-social/indigo/api/atproto" "github.com/bluesky-social/indigo/xrpc" - "github.com/ipfs/go-cid" "github.com/labstack/echo/v4" "github.com/streamplace/oatproxy/pkg/oatproxy" "go.opentelemetry.io/otel" @@ -41,7 +39,7 @@ func resolveRepoService(ctx context.Context, repo string) (string, string, strin var maxBlobSize int64 = 1024 * 1024 * 10 // 10MB -func (s *Server) handleComAtprotoRepoUploadBlob(ctx context.Context, r io.Reader, contentType string) (*comatprototypes.RepoUploadBlob_Output, error) { +func (s *Server) handleComAtprotoRepoUploadBlob(ctx context.Context, r io.Reader, contentType string) (*comatproto.RepoUploadBlob_Output, error) { ctx, span := otel.Tracer("server").Start(ctx, "handleComAtprotoRepoUploadBlob") defer span.End() @@ -61,7 +59,7 @@ func (s *Server) handleComAtprotoRepoUploadBlob(ctx context.Context, r io.Reader return nil, echo.NewHTTPError(http.StatusInternalServerError, "failed to copy reader to buffer") } - var out comatprototypes.RepoUploadBlob_Output + var out comatproto.RepoUploadBlob_Output err = client.Do(ctx, xrpc.Procedure, contentType, "com.atproto.repo.uploadBlob", nil, bytes.NewReader(buf.Bytes()), &out) @@ -73,13 +71,13 @@ func (s *Server) handleComAtprotoRepoUploadBlob(ctx context.Context, r io.Reader return &out, nil } -func (s *Server) handleComAtprotoRepoDescribeRepo(ctx context.Context, repo string) (*comatprototypes.RepoDescribeRepo_Output, error) { +func (s *Server) handleComAtprotoRepoDescribeRepo(ctx context.Context, repo string) (*comatproto.RepoDescribeRepo_Output, error) { isLocal, svc, err := s.isLocalPDS(ctx, repo) if err != nil { return nil, fmt.Errorf("error checking for local PDS: %w", err) } if !isLocal { - var out comatprototypes.RepoDescribeRepo_Output + var out comatproto.RepoDescribeRepo_Output params := make(map[string]interface{}) params["repo"] = repo @@ -92,7 +90,7 @@ func (s *Server) handleComAtprotoRepoDescribeRepo(ctx context.Context, repo stri } - return &comatprototypes.RepoDescribeRepo_Output{ + return &comatproto.RepoDescribeRepo_Output{ Handle: s.cli.MyDID(), Did: s.cli.MyDID(), DidDoc: atproto.DIDDoc(s.cli.BroadcasterHost), @@ -103,13 +101,13 @@ func (s *Server) handleComAtprotoRepoDescribeRepo(ctx context.Context, repo stri }, nil } -func (s *Server) handleComAtprotoRepoListRecords(ctx context.Context, collection string, cursor string, limit int, repo string, reverse *bool) (*comatprototypes.RepoListRecords_Output, error) { +func (s *Server) handleComAtprotoRepoListRecords(ctx context.Context, collection string, cursor string, limit int, repo string, reverse *bool) (*comatproto.RepoListRecords_Output, error) { isLocal, svc, err := s.isLocalPDS(ctx, repo) if err != nil { return nil, fmt.Errorf("error checking for local PDS: %w", err) } if !isLocal { - var out comatprototypes.RepoListRecords_Output + var out comatproto.RepoListRecords_Output params := make(map[string]interface{}) params["collection"] = collection if cursor != "" { @@ -131,39 +129,16 @@ func (s *Server) handleComAtprotoRepoListRecords(ctx context.Context, collection return &out, nil } - r, ses, err := atproto.OpenLexiconRepo(ctx) - if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoListRecords: failed to open repo: %w", err) - } - out := &comatprototypes.RepoListRecords_Output{ - Records: []*comatprototypes.RepoListRecords_Record{}, - } - err = r.ForEach(ctx, "", func(rkey string, c cid.Cid) error { - val, err := atproto.GetRecordCBOR(ctx, ses, c, collection, rkey) - if err != nil { - return fmt.Errorf("handleComAtprotoRepoListRecords: failed to get record for collection %q, rkey %q: %w", collection, rkey, err) - } - out.Records = append(out.Records, &comatprototypes.RepoListRecords_Record{ - Uri: fmt.Sprintf("at://%s/%s/%s", repo, collection, rkey), - Cid: c.String(), - Value: &lexutil.LexiconTypeDecoder{Val: val}, - }) - - return nil - }) - if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoListRecords: error iterating records for collection %q: %w", collection, err) - } - return out, nil + return atproto.LexiconRepoListRecords(ctx, collection, cursor, limit, repo, reverse) } -func (s *Server) handleComAtprotoRepoGetRecord(ctx context.Context, c string, collection string, repo string, rkey string) (*comatprototypes.RepoGetRecord_Output, error) { +func (s *Server) handleComAtprotoRepoGetRecord(ctx context.Context, c string, collection string, repo string, rkey string) (*comatproto.RepoGetRecord_Output, error) { isLocal, svc, err := s.isLocalPDS(ctx, repo) if err != nil { return nil, fmt.Errorf("error checking for local PDS: %w", err) } if !isLocal { - var out comatprototypes.RepoGetRecord_Output + var out comatproto.RepoGetRecord_Output params := make(map[string]interface{}) params["repo"] = repo params["collection"] = collection @@ -180,22 +155,5 @@ func (s *Server) handleComAtprotoRepoGetRecord(ctx context.Context, c string, co return &out, nil } - r, ses, err := atproto.OpenLexiconRepo(ctx) - if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to open repo: %w", err) - } - outCID, _, err := r.GetRecord(ctx, fmt.Sprintf("%s/%s", collection, rkey)) - if err != nil { - return nil, err - } - rec, err := atproto.GetRecordCBOR(ctx, ses, outCID, collection, rkey) - if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to get record: %w", err) - } - str := outCID.String() - return &comatprototypes.RepoGetRecord_Output{ - Uri: fmt.Sprintf("at://%s/%s/%s", repo, collection, rkey), - Cid: &str, - Value: &lexutil.LexiconTypeDecoder{Val: rec}, - }, nil + return atproto.LexiconRepoGetRecord(ctx, repo, collection, rkey) } diff --git a/pkg/spxrpc/com_atproto_sync.go b/pkg/spxrpc/com_atproto_sync.go index cd5a69d35..9c822b525 100644 --- a/pkg/spxrpc/com_atproto_sync.go +++ b/pkg/spxrpc/com_atproto_sync.go @@ -9,14 +9,8 @@ import ( "strconv" comatprototypes "github.com/bluesky-social/indigo/api/atproto" - "github.com/bluesky-social/indigo/carstore" "github.com/bluesky-social/indigo/events" - "github.com/bluesky-social/indigo/repo" - "github.com/bluesky-social/indigo/util" "github.com/gorilla/websocket" - "github.com/ipfs/go-cid" - cbor "github.com/ipfs/go-ipld-cbor" - "github.com/ipld/go-car" "github.com/labstack/echo/v4" "stream.place/streamplace/pkg/atproto" "stream.place/streamplace/pkg/log" @@ -37,52 +31,11 @@ func (s *Server) handleComAtprotoSyncListRepos(ctx context.Context, cursor strin } func (s *Server) handleComAtprotoSyncGetRecord(ctx context.Context, collection string, did string, rkey string) (io.Reader, error) { - _, robs, err := atproto.OpenLexiconRepo(ctx) + bs, err := atproto.LexiconRepoMerkleProof(ctx, collection, rkey) if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to open repo: %w", err) - } - - bs := util.NewLoggingBstore(robs) - - root, err := atproto.CarStore.GetUserRepoHead(ctx, atproto.RepoUser) - if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to get user repo head: %w", err) - } - - log.Warn(ctx, "got root", "root", root.String()) - - r, err := repo.OpenRepo(ctx, bs, root) - if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to open repo: %w", err) - } - - _, _, err = r.GetRecordBytes(ctx, collection+"/"+rkey) - if err != nil { - return nil, fmt.Errorf("handleComAtprotoRepoGetRecord: failed to get record bytes: %w", err) - } - - blocks := bs.GetLoggedBlocks() - - buf := new(bytes.Buffer) - hb, err := cbor.DumpObject(&car.CarHeader{ - Roots: []cid.Cid{root}, - Version: 1, - }) - if err != nil { - return nil, fmt.Errorf("failed to dump car header: %w", err) - } - if _, err := carstore.LdWrite(buf, hb); err != nil { return nil, err } - - for _, blk := range blocks { - log.Warn(ctx, "writing block", "cid", blk.Cid().String(), "version", blk.Cid().Version()) - if _, err := carstore.LdWrite(buf, blk.Cid().Bytes(), blk.RawData()); err != nil { - return nil, err - } - } - - return bytes.NewReader(buf.Bytes()), nil + return bytes.NewReader(bs), nil } var upgrader = websocket.Upgrader{