diff --git a/pkg/model/follow.go b/pkg/model/follow.go index 796889605..d911d239b 100644 --- a/pkg/model/follow.go +++ b/pkg/model/follow.go @@ -49,3 +49,12 @@ func (m *DBModel) GetUserFollowers(ctx context.Context, userDID string) ([]Follo var follows []Follow return follows, m.DB.Where("subject_did = ?", userDID).Find(&follows).Error } + +func (m *DBModel) GetUserFollowingUser(ctx context.Context, userDID, subjectDID string) (*Follow, error) { + var follow Follow + result := m.DB.Where("user_did = ? AND subject_did = ?", userDID, subjectDID).First(&follow) + if result.RowsAffected == 0 { + return nil, nil + } + return &follow, result.Error +} diff --git a/pkg/model/model.go b/pkg/model/model.go index 119572d9b..d2d2961a7 100644 --- a/pkg/model/model.go +++ b/pkg/model/model.go @@ -54,6 +54,7 @@ type Model interface { CreateFollow(ctx context.Context, userDID, rev string, follow *bsky.GraphFollow) error GetUserFollowing(ctx context.Context, userDID string) ([]Follow, error) GetUserFollowers(ctx context.Context, userDID string) ([]Follow, error) + GetUserFollowingUser(ctx context.Context, userDID, subjectDID string) (*Follow, error) DeleteFollow(ctx context.Context, userDID, rev string) error GetFollowersNotificationTokens(userDID string) ([]string, error) diff --git a/pkg/spxrpc/graph.go b/pkg/spxrpc/graph.go index 4ceece02a..7c7ee71f8 100644 --- a/pkg/spxrpc/graph.go +++ b/pkg/spxrpc/graph.go @@ -5,8 +5,8 @@ import ( "fmt" "github.com/bluesky-social/indigo/api/atproto" + "github.com/bluesky-social/indigo/atproto/syntax" "go.opentelemetry.io/otel" - "stream.place/streamplace/pkg/log" placestreamtypes "stream.place/streamplace/pkg/streamplace" ) @@ -14,33 +14,23 @@ func (s *Server) handlePlaceStreamGraphGetFollowingUser(ctx context.Context, use ctx, span := otel.Tracer("server").Start(ctx, "handlePlaceStreamGraphGetFollowingUser") defer span.End() - if userDID == "" || !isValidDID(userDID) { - log.Error(ctx, "Missing or invalid user DID") - return &placestreamtypes.GraphGetFollowingUser_Output{}, nil + _, didErr := syntax.ParseDID(userDID) + if userDID == "" || didErr != nil { + return nil, fmt.Errorf("Missing or invalid user DID") } - follows, err := s.model.GetUserFollowing(ctx, userDID) + follow, err := s.model.GetUserFollowingUser(ctx, userDID, subjectDID) if err != nil { - log.Error(ctx, "Failed to get user following", "error", err) - return &placestreamtypes.GraphGetFollowingUser_Output{}, nil + return nil, fmt.Errorf("Failed to get user following: %w", err) } - for _, follow := range follows { - if follow.SubjectDID == subjectDID { - // User is following the subject, return the follow reference - return &placestreamtypes.GraphGetFollowingUser_Output{ - Follow: &atproto.RepoStrongRef{ - Cid: "", // We don't store CID in our model - Uri: fmt.Sprintf("at://%s/app.bsky.graph.follow/%s", userDID, follow.RKey), - }, - }, nil + output := &placestreamtypes.GraphGetFollowingUser_Output{} + if follow != nil { + output.Follow = &atproto.RepoStrongRef{ + Cid: "", // We don't store CID in our model + Uri: fmt.Sprintf("at://%s/app.bsky.graph.follow/%s", userDID, follow.RKey), } } - // User is not following the subject - return &placestreamtypes.GraphGetFollowingUser_Output{}, nil -} - -func isValidDID(did string) bool { - return len(did) > 0 && (did[:7] == "did:plc" || did[:7] == "did:web") + return output, nil }