diff --git a/internal/api/handlers/community/get.go b/internal/api/handlers/community/get.go index 42cfdf7..103c22b 100644 --- a/internal/api/handlers/community/get.go +++ b/internal/api/handlers/community/get.go @@ -5,18 +5,21 @@ import ( "log" "net/http" + "Coves/internal/api/handlers/common" "Coves/internal/core/communities" ) // GetHandler handles community retrieval type GetHandler struct { service communities.Service + repo communities.Repository } // NewGetHandler creates a new get handler -func NewGetHandler(service communities.Service) *GetHandler { +func NewGetHandler(service communities.Service, repo communities.Repository) *GetHandler { return &GetHandler{ service: service, + repo: repo, } } @@ -42,6 +45,9 @@ func (h *GetHandler) HandleGet(w http.ResponseWriter, r *http.Request) { return } + // Populate viewer state (viewer.subscribed) if authenticated + common.PopulateCommunityViewerState(r.Context(), r, h.repo, []*communities.Community{community}) + // Convert to detailed view for API response view := community.ToCommunityViewDetailed() diff --git a/internal/api/routes/community.go b/internal/api/routes/community.go index 9479a3f..a338a7e 100644 --- a/internal/api/routes/community.go +++ b/internal/api/routes/community.go @@ -14,7 +14,7 @@ import ( func RegisterCommunityRoutes(r chi.Router, service communities.Service, repo communities.Repository, authMiddleware *middleware.OAuthAuthMiddleware, allowedCommunityCreators []string) { // Initialize handlers createHandler := community.NewCreateHandler(service, allowedCommunityCreators) - getHandler := community.NewGetHandler(service) + getHandler := community.NewGetHandler(service, repo) updateHandler := community.NewUpdateHandler(service) listHandler := community.NewListHandler(service, repo) searchHandler := community.NewSearchHandler(service) @@ -23,7 +23,8 @@ func RegisterCommunityRoutes(r chi.Router, service communities.Service, repo com // Query endpoints (GET) - public access, optional auth for viewer state // social.coves.community.get - get a single community by identifier - r.Get("/xrpc/social.coves.community.get", getHandler.HandleGet) + // Uses OptionalAuth to populate viewer.subscribed when authenticated + r.With(authMiddleware.OptionalAuth).Get("/xrpc/social.coves.community.get", getHandler.HandleGet) // social.coves.community.list - list communities with filters // Uses OptionalAuth to populate viewer.subscribed when authenticated diff --git a/tests/integration/community_get_viewer_state_test.go b/tests/integration/community_get_viewer_state_test.go new file mode 100644 index 0000000..94481ed --- /dev/null +++ b/tests/integration/community_get_viewer_state_test.go @@ -0,0 +1,151 @@ +package integration + +import ( + "Coves/internal/api/handlers/community" + "Coves/internal/api/middleware" + "Coves/internal/core/communities" + "Coves/internal/db/postgres" + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/go-chi/chi/v5" +) + +// getViewerMockService reuses mockCommunityService but implements GetCommunity +// against the real repository, since the get endpoint is what's under test. +type getViewerMockService struct { + mockCommunityService +} + +func (m *getViewerMockService) GetCommunity(ctx context.Context, identifier string) (*communities.Community, error) { + return m.repo.GetByDID(ctx, identifier) +} + +// TestCommunityGet_ViewerState tests that the get community endpoint +// populates viewer.subscribed for authenticated users, matching the +// social.coves.community.get lexicon promise ("viewer state will be +// included if authenticated"). +func TestCommunityGet_ViewerState(t *testing.T) { + db := setupTestDB(t) + defer func() { + if err := db.Close(); err != nil { + t.Logf("Failed to close database: %v", err) + } + }() + + repo := postgres.NewCommunityRepository(db) + ctx := context.Background() + + // Create two communities: the user subscribes to the first only + baseSuffix := time.Now().UnixNano() + communityDIDs := make([]string, 2) + for i := 0; i < 2; i++ { + communityDID := generateTestDID(fmt.Sprintf("%d%d", baseSuffix, i)) + communityDIDs[i] = communityDID + comm := &communities.Community{ + DID: communityDID, + Handle: fmt.Sprintf("c-getviewer-%d-%d.coves.local", baseSuffix, i), + Name: fmt.Sprintf("getviewer-test-%d", i), + DisplayName: fmt.Sprintf("Get Viewer Test Community %d", i), + OwnerDID: "did:web:coves.local", + CreatedByDID: "did:plc:testcreator", + HostedByDID: "did:web:coves.local", + Visibility: "public", + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + } + if _, err := repo.Create(ctx, comm); err != nil { + t.Fatalf("Failed to create community %d: %v", i, err) + } + } + + testUserDID := fmt.Sprintf("did:plc:getviewertestuser%d", baseSuffix) + sub := &communities.Subscription{ + UserDID: testUserDID, + CommunityDID: communityDIDs[0], + ContentVisibility: 3, + SubscribedAt: time.Now(), + } + if _, err := repo.Subscribe(ctx, sub); err != nil { + t.Fatalf("Failed to subscribe to community 0: %v", err) + } + + mockService := &getViewerMockService{mockCommunityService{repo: repo}} + getHandler := community.NewGetHandler(mockService, repo) + + type getResponse struct { + DID string `json:"did"` + Viewer *struct { + Subscribed *bool `json:"subscribed"` + } `json:"viewer"` + } + + doGet := func(t *testing.T, r chi.Router, communityDID string) getResponse { + t.Helper() + req := httptest.NewRequest("GET", "/xrpc/social.coves.community.get?community="+communityDID, nil) + rec := httptest.NewRecorder() + r.ServeHTTP(rec, req) + + if rec.Code != http.StatusOK { + t.Fatalf("Expected status 200, got %d: %s", rec.Code, rec.Body.String()) + } + var resp getResponse + if err := json.NewDecoder(rec.Body).Decode(&resp); err != nil { + t.Fatalf("Failed to decode response: %v", err) + } + return resp + } + + t.Run("authenticated subscriber sees viewer.subscribed=true", func(t *testing.T) { + r := chi.NewRouter() + r.Use(func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + ctx := middleware.SetTestUserDID(req.Context(), testUserDID) + next.ServeHTTP(w, req.WithContext(ctx)) + }) + }) + r.Get("/xrpc/social.coves.community.get", getHandler.HandleGet) + + resp := doGet(t, r, communityDIDs[0]) + if resp.Viewer == nil || resp.Viewer.Subscribed == nil { + t.Fatalf("Expected populated viewer.subscribed, got viewer=%+v", resp.Viewer) + } + if !*resp.Viewer.Subscribed { + t.Errorf("Expected viewer.subscribed=true for subscribed community %s", communityDIDs[0]) + } + }) + + t.Run("authenticated non-subscriber sees viewer.subscribed=false", func(t *testing.T) { + r := chi.NewRouter() + r.Use(func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + ctx := middleware.SetTestUserDID(req.Context(), testUserDID) + next.ServeHTTP(w, req.WithContext(ctx)) + }) + }) + r.Get("/xrpc/social.coves.community.get", getHandler.HandleGet) + + resp := doGet(t, r, communityDIDs[1]) + if resp.Viewer == nil || resp.Viewer.Subscribed == nil { + t.Fatalf("Expected populated viewer.subscribed, got viewer=%+v", resp.Viewer) + } + if *resp.Viewer.Subscribed { + t.Errorf("Expected viewer.subscribed=false for unsubscribed community %s", communityDIDs[1]) + } + }) + + t.Run("unauthenticated request has nil viewer state", func(t *testing.T) { + r := chi.NewRouter() + r.Get("/xrpc/social.coves.community.get", getHandler.HandleGet) + + resp := doGet(t, r, communityDIDs[0]) + if resp.Viewer != nil { + t.Errorf("Expected nil viewer for unauthenticated request, got %+v", resp.Viewer) + } + }) +}