diff --git a/cmd/server/main.go b/cmd/server/main.go index b4af6f3..5ad5be0 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -276,8 +276,9 @@ func main() { log.Printf(" - Communities will be created at: %s", defaultPDS) log.Printf(" - PDS will generate and manage all DIDs and keys") - // Initialize community service (no longer needs didGenerator directly) - communityService := communities.NewCommunityService(communityRepo, defaultPDS, instanceDID, instanceDomain, provisioner) + // Initialize community service with OAuth client for user DPoP authentication + // OAuth client is required for subscribe/unsubscribe/block/unblock operations + communityService := communities.NewCommunityService(communityRepo, defaultPDS, instanceDID, instanceDomain, provisioner, oauthClient, oauthStore) // Authenticate Coves instance with PDS to enable community record writes // The instance needs a PDS account to write community records it owns diff --git a/internal/api/handlers/community/block.go b/internal/api/handlers/community/block.go index 4e6f256..e84aea0 100644 --- a/internal/api/handlers/community/block.go +++ b/internal/api/handlers/community/block.go @@ -47,38 +47,17 @@ func (h *BlockHandler) HandleBlock(w http.ResponseWriter, r *http.Request) { return } - // Extract authenticated user DID and access token from request context (injected by auth middleware) - userDID := middleware.GetUserDID(r) - if userDID == "" { + // Get OAuth session from context (injected by auth middleware) + // The session contains the user's DID and credentials needed for DPoP authentication + session := middleware.GetOAuthSession(r) + if session == nil { writeError(w, http.StatusUnauthorized, "AuthRequired", "Authentication required") return } - userAccessToken := middleware.GetUserAccessToken(r) - if userAccessToken == "" { - writeError(w, http.StatusUnauthorized, "AuthRequired", "Missing access token") - return - } - - // Resolve community identifier (handle or DID) to DID - // This allows users to block by handle: @gaming.community.coves.social or !gaming@coves.social - communityDID, err := h.service.ResolveCommunityIdentifier(r.Context(), req.Community) - if err != nil { - if communities.IsNotFound(err) { - writeError(w, http.StatusNotFound, "CommunityNotFound", "Community not found") - return - } - if communities.IsValidationError(err) { - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - return - } - log.Printf("Failed to resolve community identifier %s: %v", req.Community, err) - writeError(w, http.StatusInternalServerError, "InternalError", "Failed to resolve community") - return - } - - // Block via service (write-forward to PDS) using resolved DID - block, err := h.service.BlockCommunity(r.Context(), userDID, userAccessToken, communityDID) + // Block via service (write-forward to PDS with DPoP authentication) + // Service handles identifier resolution (DIDs, handles, scoped identifiers) + block, err := h.service.BlockCommunity(r.Context(), session, req.Community) if err != nil { handleServiceError(w, err) return @@ -125,38 +104,17 @@ func (h *BlockHandler) HandleUnblock(w http.ResponseWriter, r *http.Request) { return } - // Extract authenticated user DID and access token from request context (injected by auth middleware) - userDID := middleware.GetUserDID(r) - if userDID == "" { + // Get OAuth session from context (injected by auth middleware) + // The session contains the user's DID and credentials needed for DPoP authentication + session := middleware.GetOAuthSession(r) + if session == nil { writeError(w, http.StatusUnauthorized, "AuthRequired", "Authentication required") return } - userAccessToken := middleware.GetUserAccessToken(r) - if userAccessToken == "" { - writeError(w, http.StatusUnauthorized, "AuthRequired", "Missing access token") - return - } - - // Resolve community identifier (handle or DID) to DID - // This allows users to unblock by handle: @gaming.community.coves.social or !gaming@coves.social - communityDID, err := h.service.ResolveCommunityIdentifier(r.Context(), req.Community) - if err != nil { - if communities.IsNotFound(err) { - writeError(w, http.StatusNotFound, "CommunityNotFound", "Community not found") - return - } - if communities.IsValidationError(err) { - writeError(w, http.StatusBadRequest, "InvalidRequest", err.Error()) - return - } - log.Printf("Failed to resolve community identifier %s: %v", req.Community, err) - writeError(w, http.StatusInternalServerError, "InternalError", "Failed to resolve community") - return - } - - // Unblock via service (delete record on PDS) using resolved DID - err = h.service.UnblockCommunity(r.Context(), userDID, userAccessToken, communityDID) + // Unblock via service (delete record on PDS with DPoP authentication) + // Service handles identifier resolution (DIDs, handles, scoped identifiers) + err := h.service.UnblockCommunity(r.Context(), session, req.Community) if err != nil { handleServiceError(w, err) return diff --git a/internal/api/handlers/community/create_test.go b/internal/api/handlers/community/create_test.go index 205d77b..709e16e 100644 --- a/internal/api/handlers/community/create_test.go +++ b/internal/api/handlers/community/create_test.go @@ -10,6 +10,8 @@ import ( "net/http/httptest" "testing" "time" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" ) // mockCommunityService implements communities.Service for testing @@ -49,11 +51,11 @@ func (m *mockCommunityService) SearchCommunities(ctx context.Context, req commun return nil, 0, nil } -func (m *mockCommunityService) SubscribeToCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { +func (m *mockCommunityService) SubscribeToCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { return nil, nil } -func (m *mockCommunityService) UnsubscribeFromCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string) error { +func (m *mockCommunityService) UnsubscribeFromCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { return nil } @@ -65,11 +67,11 @@ func (m *mockCommunityService) GetCommunitySubscribers(ctx context.Context, comm return nil, nil } -func (m *mockCommunityService) BlockCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string) (*communities.CommunityBlock, error) { +func (m *mockCommunityService) BlockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) (*communities.CommunityBlock, error) { return nil, nil } -func (m *mockCommunityService) UnblockCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string) error { +func (m *mockCommunityService) UnblockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { return nil } diff --git a/internal/api/handlers/community/subscribe.go b/internal/api/handlers/community/subscribe.go index 0f24548..0283493 100644 --- a/internal/api/handlers/community/subscribe.go +++ b/internal/api/handlers/community/subscribe.go @@ -51,23 +51,17 @@ func (h *SubscribeHandler) HandleSubscribe(w http.ResponseWriter, r *http.Reques return } - // Extract authenticated user DID and access token from request context (injected by auth middleware) - // Note: contentVisibility defaults and clamping handled by service layer - userDID := middleware.GetUserDID(r) - if userDID == "" { + // Get OAuth session from context (injected by auth middleware) + // The session contains the user's DID and credentials needed for DPoP authentication + session := middleware.GetOAuthSession(r) + if session == nil { writeError(w, http.StatusUnauthorized, "AuthRequired", "Authentication required") return } - userAccessToken := middleware.GetUserAccessToken(r) - if userAccessToken == "" { - writeError(w, http.StatusUnauthorized, "AuthRequired", "Missing access token") - return - } - - // Subscribe via service (write-forward to PDS) + // Subscribe via service (write-forward to PDS with DPoP authentication) // Service handles identifier resolution (DIDs, handles, scoped identifiers) - subscription, err := h.service.SubscribeToCommunity(r.Context(), userDID, userAccessToken, req.Community, req.ContentVisibility) + subscription, err := h.service.SubscribeToCommunity(r.Context(), session, req.Community, req.ContentVisibility) if err != nil { handleServiceError(w, err) return @@ -117,22 +111,17 @@ func (h *SubscribeHandler) HandleUnsubscribe(w http.ResponseWriter, r *http.Requ return } - // Extract authenticated user DID and access token from request context (injected by auth middleware) - userDID := middleware.GetUserDID(r) - if userDID == "" { + // Get OAuth session from context (injected by auth middleware) + // The session contains the user's DID and credentials needed for DPoP authentication + session := middleware.GetOAuthSession(r) + if session == nil { writeError(w, http.StatusUnauthorized, "AuthRequired", "Authentication required") return } - userAccessToken := middleware.GetUserAccessToken(r) - if userAccessToken == "" { - writeError(w, http.StatusUnauthorized, "AuthRequired", "Missing access token") - return - } - - // Unsubscribe via service (delete record on PDS) + // Unsubscribe via service (delete record on PDS with DPoP authentication) // Service handles identifier resolution (DIDs, handles, scoped identifiers) - err := h.service.UnsubscribeFromCommunity(r.Context(), userDID, userAccessToken, req.Community) + err := h.service.UnsubscribeFromCommunity(r.Context(), session, req.Community) if err != nil { handleServiceError(w, err) return diff --git a/internal/api/handlers/community/subscribe_test.go b/internal/api/handlers/community/subscribe_test.go index 54fd4a6..b5abdcc 100644 --- a/internal/api/handlers/community/subscribe_test.go +++ b/internal/api/handlers/community/subscribe_test.go @@ -11,12 +11,26 @@ import ( "net/http/httptest" "testing" "time" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" ) +// createTestOAuthSession creates a mock OAuth session for testing +func createTestOAuthSession(did string) *oauth.ClientSessionData { + parsedDID, _ := syntax.ParseDID(did) + return &oauth.ClientSessionData{ + AccountDID: parsedDID, + SessionID: "test-session", + HostURL: "http://localhost:3001", + AccessToken: "test-access-token", + } +} + // subscribeTestService implements communities.Service for subscribe handler tests type subscribeTestService struct { - subscribeFunc func(ctx context.Context, userDID, accessToken, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) - unsubscribeFunc func(ctx context.Context, userDID, accessToken, communityIdentifier string) error + subscribeFunc func(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) + unsubscribeFunc func(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error } func (m *subscribeTestService) CreateCommunity(ctx context.Context, req communities.CreateCommunityRequest) (*communities.Community, error) { @@ -39,9 +53,13 @@ func (m *subscribeTestService) SearchCommunities(ctx context.Context, req commun return nil, 0, nil } -func (m *subscribeTestService) SubscribeToCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { +func (m *subscribeTestService) SubscribeToCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { if m.subscribeFunc != nil { - return m.subscribeFunc(ctx, userDID, accessToken, communityIdentifier, contentVisibility) + return m.subscribeFunc(ctx, session, communityIdentifier, contentVisibility) + } + userDID := "" + if session != nil { + userDID = session.AccountDID.String() } return &communities.Subscription{ UserDID: userDID, @@ -52,9 +70,9 @@ func (m *subscribeTestService) SubscribeToCommunity(ctx context.Context, userDID }, nil } -func (m *subscribeTestService) UnsubscribeFromCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string) error { +func (m *subscribeTestService) UnsubscribeFromCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { if m.unsubscribeFunc != nil { - return m.unsubscribeFunc(ctx, userDID, accessToken, communityIdentifier) + return m.unsubscribeFunc(ctx, session, communityIdentifier) } return nil } @@ -67,11 +85,11 @@ func (m *subscribeTestService) GetCommunitySubscribers(ctx context.Context, comm return nil, nil } -func (m *subscribeTestService) BlockCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string) (*communities.CommunityBlock, error) { +func (m *subscribeTestService) BlockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) (*communities.CommunityBlock, error) { return nil, nil } -func (m *subscribeTestService) UnblockCommunity(ctx context.Context, userDID, accessToken, communityIdentifier string) error { +func (m *subscribeTestService) UnblockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { return nil } @@ -144,8 +162,12 @@ func TestSubscribeHandler_Subscribe_Success(t *testing.T) { t.Run(tc.name, func(t *testing.T) { var receivedIdentifier string mockService := &subscribeTestService{ - subscribeFunc: func(ctx context.Context, userDID, accessToken, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { + subscribeFunc: func(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { receivedIdentifier = communityIdentifier + userDID := "" + if session != nil { + userDID = session.AccountDID.String() + } return &communities.Subscription{ UserDID: userDID, CommunityDID: "did:plc:resolved", @@ -167,9 +189,9 @@ func TestSubscribeHandler_Subscribe_Success(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.subscribe", bytes.NewBuffer(bodyBytes)) req.Header.Set("Content-Type", "application/json") - // Inject auth context - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:testuser") - ctx = context.WithValue(ctx, middleware.UserAccessToken, "test-token") + // Inject OAuth session into context + session := createTestOAuthSession("did:plc:testuser") + ctx := context.WithValue(req.Context(), middleware.OAuthSessionKey, session) req = req.WithContext(ctx) w := httptest.NewRecorder() @@ -244,8 +266,8 @@ func TestSubscribeHandler_Subscribe_RequiresCommunity(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.subscribe", bytes.NewBuffer(bodyBytes)) req.Header.Set("Content-Type", "application/json") - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:testuser") - ctx = context.WithValue(ctx, middleware.UserAccessToken, "test-token") + session := createTestOAuthSession("did:plc:testuser") + ctx := context.WithValue(req.Context(), middleware.OAuthSessionKey, session) req = req.WithContext(ctx) w := httptest.NewRecorder() @@ -286,7 +308,7 @@ func TestSubscribeHandler_Subscribe_ServiceErrors(t *testing.T) { for _, tc := range tests { t.Run(tc.name, func(t *testing.T) { mockService := &subscribeTestService{ - subscribeFunc: func(ctx context.Context, userDID, accessToken, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { + subscribeFunc: func(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string, contentVisibility int) (*communities.Subscription, error) { return nil, tc.serviceErr }, } @@ -302,8 +324,8 @@ func TestSubscribeHandler_Subscribe_ServiceErrors(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.subscribe", bytes.NewBuffer(bodyBytes)) req.Header.Set("Content-Type", "application/json") - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:testuser") - ctx = context.WithValue(ctx, middleware.UserAccessToken, "test-token") + session := createTestOAuthSession("did:plc:testuser") + ctx := context.WithValue(req.Context(), middleware.OAuthSessionKey, session) req = req.WithContext(ctx) w := httptest.NewRecorder() @@ -353,7 +375,7 @@ func TestSubscribeHandler_Unsubscribe_Success(t *testing.T) { t.Run(tc.name, func(t *testing.T) { var receivedIdentifier string mockService := &subscribeTestService{ - unsubscribeFunc: func(ctx context.Context, userDID, accessToken, communityIdentifier string) error { + unsubscribeFunc: func(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { receivedIdentifier = communityIdentifier return nil }, @@ -369,8 +391,8 @@ func TestSubscribeHandler_Unsubscribe_Success(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.unsubscribe", bytes.NewBuffer(bodyBytes)) req.Header.Set("Content-Type", "application/json") - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:testuser") - ctx = context.WithValue(ctx, middleware.UserAccessToken, "test-token") + session := createTestOAuthSession("did:plc:testuser") + ctx := context.WithValue(req.Context(), middleware.OAuthSessionKey, session) req = req.WithContext(ctx) w := httptest.NewRecorder() @@ -399,7 +421,7 @@ func TestSubscribeHandler_Unsubscribe_Success(t *testing.T) { func TestSubscribeHandler_Unsubscribe_SubscriptionNotFound(t *testing.T) { mockService := &subscribeTestService{ - unsubscribeFunc: func(ctx context.Context, userDID, accessToken, communityIdentifier string) error { + unsubscribeFunc: func(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { return communities.ErrSubscriptionNotFound }, } @@ -414,8 +436,8 @@ func TestSubscribeHandler_Unsubscribe_SubscriptionNotFound(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.unsubscribe", bytes.NewBuffer(bodyBytes)) req.Header.Set("Content-Type", "application/json") - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:testuser") - ctx = context.WithValue(ctx, middleware.UserAccessToken, "test-token") + session := createTestOAuthSession("did:plc:testuser") + ctx := context.WithValue(req.Context(), middleware.OAuthSessionKey, session) req = req.WithContext(ctx) w := httptest.NewRecorder() @@ -456,8 +478,8 @@ func TestSubscribeHandler_InvalidJSON(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.subscribe", bytes.NewBufferString("invalid json")) req.Header.Set("Content-Type", "application/json") - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:testuser") - ctx = context.WithValue(ctx, middleware.UserAccessToken, "test-token") + session := createTestOAuthSession("did:plc:testuser") + ctx := context.WithValue(req.Context(), middleware.OAuthSessionKey, session) req = req.WithContext(ctx) w := httptest.NewRecorder() @@ -468,7 +490,7 @@ func TestSubscribeHandler_InvalidJSON(t *testing.T) { } } -func TestSubscribeHandler_RequiresAccessToken(t *testing.T) { +func TestSubscribeHandler_RequiresOAuthSession(t *testing.T) { mockService := &subscribeTestService{} handler := NewSubscribeHandler(mockService) @@ -480,9 +502,7 @@ func TestSubscribeHandler_RequiresAccessToken(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.subscribe", bytes.NewBuffer(bodyBytes)) req.Header.Set("Content-Type", "application/json") - // User DID but no access token - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:testuser") - req = req.WithContext(ctx) + // No OAuth session in context w := httptest.NewRecorder() handler.HandleSubscribe(w, req) diff --git a/internal/atproto/pds/errors.go b/internal/atproto/pds/errors.go index 003df74..52d6920 100644 --- a/internal/atproto/pds/errors.go +++ b/internal/atproto/pds/errors.go @@ -27,3 +27,8 @@ var ( func IsAuthError(err error) bool { return errors.Is(err, ErrUnauthorized) || errors.Is(err, ErrForbidden) } + +// IsConflictError returns true if the error indicates a conflict (e.g., duplicate record). +func IsConflictError(err error) bool { + return errors.Is(err, ErrConflict) +} diff --git a/internal/core/communities/interfaces.go b/internal/core/communities/interfaces.go index 84b3055..fe2acbc 100644 --- a/internal/core/communities/interfaces.go +++ b/internal/core/communities/interfaces.go @@ -1,6 +1,10 @@ package communities -import "context" +import ( + "context" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" +) // Repository defines the interface for community data persistence // This is the AppView's indexed view of communities from the firehose @@ -66,14 +70,16 @@ type Service interface { SearchCommunities(ctx context.Context, req SearchCommunitiesRequest) ([]*Community, int, error) // Subscription operations (write-forward: creates record in user's PDS) - SubscribeToCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string, contentVisibility int) (*Subscription, error) - UnsubscribeFromCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string) error + // OAuth session is passed for DPoP authentication to the user's PDS + SubscribeToCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string, contentVisibility int) (*Subscription, error) + UnsubscribeFromCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error GetUserSubscriptions(ctx context.Context, userDID string, limit, offset int) ([]*Subscription, error) GetCommunitySubscribers(ctx context.Context, communityIdentifier string, limit, offset int) ([]*Subscription, error) // Block operations (write-forward: creates record in user's PDS) - BlockCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string) (*CommunityBlock, error) - UnblockCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string) error + // OAuth session is passed for DPoP authentication to the user's PDS + BlockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) (*CommunityBlock, error) + UnblockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error GetBlockedCommunities(ctx context.Context, userDID string, limit, offset int) ([]*CommunityBlock, error) IsBlocked(ctx context.Context, userDID, communityIdentifier string) (bool, error) diff --git a/internal/core/communities/service.go b/internal/core/communities/service.go index c4b0ac6..727693d 100644 --- a/internal/core/communities/service.go +++ b/internal/core/communities/service.go @@ -1,6 +1,8 @@ package communities import ( + oauthclient "Coves/internal/atproto/oauth" + "Coves/internal/atproto/pds" "Coves/internal/atproto/utils" "bytes" "context" @@ -14,6 +16,9 @@ import ( "strings" "sync" "time" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" ) // Community handle validation regex (DNS-valid handle: name.community.instance.com) @@ -26,11 +31,20 @@ var dnsLabelRegex = regexp.MustCompile(`^[a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0- // Domain validation (simplified - checks for valid DNS hostname structure) var domainRegex = regexp.MustCompile(`^([a-zA-Z0-9]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?\.)*[a-zA-Z]([a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?$`) +// PDSClientFactory creates PDS clients from session data. +// Used to allow injection of different auth mechanisms (OAuth for production, password for tests). +type PDSClientFactory func(ctx context.Context, session *oauth.ClientSessionData) (pds.Client, error) + type communityService struct { // Interfaces and pointers first (better alignment) repo Repository provisioner *PDSAccountProvisioner + // OAuth client/store for user PDS authentication (DPoP-based) + oauthClient *oauthclient.OAuthClient + oauthStore oauth.ClientAuthStore + pdsClientFactory PDSClientFactory // Optional, for testing. If nil, uses OAuth. + // Token refresh concurrency control // Each community gets its own mutex to prevent concurrent refresh attempts refreshMutexes map[string]*sync.Mutex @@ -52,8 +66,14 @@ const ( maxMutexCacheSize = 10000 ) -// NewCommunityService creates a new community service -func NewCommunityService(repo Repository, pdsURL, instanceDID, instanceDomain string, provisioner *PDSAccountProvisioner) Service { +// NewCommunityService creates a new community service with OAuth client for user authentication +func NewCommunityService( + repo Repository, + pdsURL, instanceDID, instanceDomain string, + provisioner *PDSAccountProvisioner, + oauthClient *oauthclient.OAuthClient, + oauthStore oauth.ClientAuthStore, +) Service { // SECURITY: Basic validation that did:web domain matches configured instanceDomain // This catches honest configuration mistakes but NOT malicious code modifications // Full verification (Phase 2) requires fetching DID document from domain @@ -74,16 +94,59 @@ func NewCommunityService(repo Repository, pdsURL, instanceDID, instanceDomain st instanceDID: instanceDID, instanceDomain: instanceDomain, provisioner: provisioner, + oauthClient: oauthClient, + oauthStore: oauthStore, refreshMutexes: make(map[string]*sync.Mutex), } } +// NewCommunityServiceWithPDSFactory creates a community service with a custom PDS client factory. +// This is primarily for testing with password-based authentication. +func NewCommunityServiceWithPDSFactory( + repo Repository, + pdsURL, instanceDID, instanceDomain string, + provisioner *PDSAccountProvisioner, + factory PDSClientFactory, +) Service { + return &communityService{ + repo: repo, + pdsURL: pdsURL, + instanceDID: instanceDID, + instanceDomain: instanceDomain, + provisioner: provisioner, + pdsClientFactory: factory, + refreshMutexes: make(map[string]*sync.Mutex), + } +} + // SetPDSAccessToken sets the PDS access token for authentication // This should be called after creating a session for the Coves instance DID on the PDS func (s *communityService) SetPDSAccessToken(token string) { s.pdsAccessToken = token } +// getPDSClient creates a PDS client from an OAuth session. +// If a custom factory was provided (for testing), uses that. +// Otherwise, uses DPoP authentication via indigo's APIClient for proper OAuth token handling. +func (s *communityService) getPDSClient(ctx context.Context, session *oauth.ClientSessionData) (pds.Client, error) { + // Use custom factory if provided (e.g., for testing with password auth) + if s.pdsClientFactory != nil { + return s.pdsClientFactory(ctx, session) + } + + // Production path: use OAuth with DPoP + if s.oauthClient == nil || s.oauthClient.ClientApp == nil { + return nil, fmt.Errorf("OAuth client not configured") + } + + client, err := pds.NewFromOAuthSession(ctx, s.oauthClient.ClientApp, session) + if err != nil { + return nil, fmt.Errorf("failed to create PDS client: %w", err) + } + + return client, nil +} + // CreateCommunity creates a new community via write-forward to PDS // V2 Flow: // 1. Service creates PDS account for community (PDS generates signing keypair) @@ -585,14 +648,14 @@ func (s *communityService) SearchCommunities(ctx context.Context, req SearchComm } // SubscribeToCommunity creates a subscription via write-forward to PDS -func (s *communityService) SubscribeToCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string, contentVisibility int) (*Subscription, error) { - if userDID == "" { - return nil, NewValidationError("userDid", "required") - } - if userAccessToken == "" { - return nil, NewValidationError("userAccessToken", "required") +// Uses OAuth session with DPoP authentication for secure PDS communication +func (s *communityService) SubscribeToCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string, contentVisibility int) (*Subscription, error) { + if session == nil { + return nil, NewValidationError("session", "required") } + userDID := session.AccountDID.String() + // Clamp contentVisibility to valid range (1-5), default to 3 if 0 or invalid if contentVisibility <= 0 || contentVisibility > 5 { contentVisibility = 3 @@ -615,6 +678,15 @@ func (s *communityService) SubscribeToCommunity(ctx context.Context, userDID, us return nil, ErrUnauthorized } + // Create PDS client for this session (DPoP authentication) + pdsClient, err := s.getPDSClient(ctx, session) + if err != nil { + return nil, fmt.Errorf("failed to create PDS client: %w", err) + } + + // Generate TID for record key + tid := syntax.NewTIDNow(0) + // Build subscription record // CRITICAL: Collection is social.coves.community.subscription (RECORD TYPE), not social.coves.community.subscribe (XRPC procedure) // This record will be created in the USER's repository: at://user_did/social.coves.community.subscription/{tid} @@ -626,10 +698,12 @@ func (s *communityService) SubscribeToCommunity(ctx context.Context, userDID, us "contentVisibility": contentVisibility, } - // Write-forward: create subscription record in user's repo using their access token - // The collection parameter refers to the record type in the repository - recordURI, recordCID, err := s.createRecordOnPDSAs(ctx, userDID, "social.coves.community.subscription", "", subRecord, userAccessToken) + // Write-forward: create subscription record in user's repo using DPoP-authenticated client + recordURI, recordCID, err := pdsClient.CreateRecord(ctx, "social.coves.community.subscription", tid.String(), subRecord) if err != nil { + if pds.IsAuthError(err) { + return nil, ErrUnauthorized + } return nil, fmt.Errorf("failed to create subscription on PDS: %w", err) } @@ -647,14 +721,14 @@ func (s *communityService) SubscribeToCommunity(ctx context.Context, userDID, us } // UnsubscribeFromCommunity removes a subscription via PDS delete -func (s *communityService) UnsubscribeFromCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string) error { - if userDID == "" { - return NewValidationError("userDid", "required") - } - if userAccessToken == "" { - return NewValidationError("userAccessToken", "required") +// Uses OAuth session with DPoP authentication for secure PDS communication +func (s *communityService) UnsubscribeFromCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { + if session == nil { + return NewValidationError("session", "required") } + userDID := session.AccountDID.String() + // Resolve community identifier communityDID, err := s.ResolveCommunityIdentifier(ctx, communityIdentifier) if err != nil { @@ -673,9 +747,18 @@ func (s *communityService) UnsubscribeFromCommunity(ctx context.Context, userDID return fmt.Errorf("invalid subscription record URI") } - // Write-forward: delete record from PDS using user's access token + // Create PDS client for this session (DPoP authentication) + pdsClient, err := s.getPDSClient(ctx, session) + if err != nil { + return fmt.Errorf("failed to create PDS client: %w", err) + } + + // Write-forward: delete record from PDS using DPoP-authenticated client // CRITICAL: Delete from social.coves.community.subscription (RECORD TYPE), not social.coves.community.unsubscribe - if err := s.deleteRecordOnPDSAs(ctx, userDID, "social.coves.community.subscription", rkey, userAccessToken); err != nil { + if err := pdsClient.DeleteRecord(ctx, "social.coves.community.subscription", rkey); err != nil { + if pds.IsAuthError(err) { + return ErrUnauthorized + } return fmt.Errorf("failed to delete subscription on PDS: %w", err) } @@ -730,20 +813,29 @@ func (s *communityService) ListCommunityMembers(ctx context.Context, communityId } // BlockCommunity blocks a community via write-forward to PDS -func (s *communityService) BlockCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string) (*CommunityBlock, error) { - if userDID == "" { - return nil, NewValidationError("userDid", "required") - } - if userAccessToken == "" { - return nil, NewValidationError("userAccessToken", "required") +// Uses OAuth session with DPoP authentication for secure PDS communication +func (s *communityService) BlockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) (*CommunityBlock, error) { + if session == nil { + return nil, NewValidationError("session", "required") } + userDID := session.AccountDID.String() + // Resolve community identifier (also verifies community exists) communityDID, err := s.ResolveCommunityIdentifier(ctx, communityIdentifier) if err != nil { return nil, err } + // Create PDS client for this session (DPoP authentication) + pdsClient, err := s.getPDSClient(ctx, session) + if err != nil { + return nil, fmt.Errorf("failed to create PDS client: %w", err) + } + + // Generate TID for record key + tid := syntax.NewTIDNow(0) + // Build block record // CRITICAL: Collection is social.coves.community.block (RECORD TYPE) // This record will be created in the USER's repository: at://user_did/social.coves.community.block/{tid} @@ -754,23 +846,20 @@ func (s *communityService) BlockCommunity(ctx context.Context, userDID, userAcce "createdAt": time.Now().Format(time.RFC3339), } - // Write-forward: create block record in user's repo using their access token + // Write-forward: create block record in user's repo using DPoP-authenticated client // Note: We don't check for existing blocks first because: // 1. The PDS may reject duplicates (depending on implementation) // 2. The repository layer handles idempotency with ON CONFLICT DO NOTHING // 3. This avoids a race condition where two concurrent requests both pass the check - recordURI, recordCID, err := s.createRecordOnPDSAs(ctx, userDID, "social.coves.community.block", "", blockRecord, userAccessToken) + recordURI, recordCID, err := pdsClient.CreateRecord(ctx, "social.coves.community.block", tid.String(), blockRecord) if err != nil { + // Check for auth errors first + if pds.IsAuthError(err) { + return nil, ErrUnauthorized + } + // Check if this is a duplicate/conflict error from PDS - // PDS should return 409 Conflict for duplicate records, but we also check common error messages - // for compatibility with different PDS implementations - errMsg := err.Error() - isDuplicate := strings.Contains(errMsg, "status 409") || // HTTP 409 Conflict - strings.Contains(errMsg, "duplicate") || - strings.Contains(errMsg, "already exists") || - strings.Contains(errMsg, "AlreadyExists") - - if isDuplicate { + if pds.IsConflictError(err) { // Fetch and return existing block from our indexed view existingBlock, getErr := s.repo.GetBlock(ctx, userDID, communityDID) if getErr == nil { @@ -804,14 +893,14 @@ func (s *communityService) BlockCommunity(ctx context.Context, userDID, userAcce } // UnblockCommunity removes a block via PDS delete -func (s *communityService) UnblockCommunity(ctx context.Context, userDID, userAccessToken, communityIdentifier string) error { - if userDID == "" { - return NewValidationError("userDid", "required") - } - if userAccessToken == "" { - return NewValidationError("userAccessToken", "required") +// Uses OAuth session with DPoP authentication for secure PDS communication +func (s *communityService) UnblockCommunity(ctx context.Context, session *oauth.ClientSessionData, communityIdentifier string) error { + if session == nil { + return NewValidationError("session", "required") } + userDID := session.AccountDID.String() + // Resolve community identifier communityDID, err := s.ResolveCommunityIdentifier(ctx, communityIdentifier) if err != nil { @@ -830,8 +919,17 @@ func (s *communityService) UnblockCommunity(ctx context.Context, userDID, userAc return fmt.Errorf("invalid block record URI") } - // Write-forward: delete record from PDS using user's access token - if err := s.deleteRecordOnPDSAs(ctx, userDID, "social.coves.community.block", rkey, userAccessToken); err != nil { + // Create PDS client for this session (DPoP authentication) + pdsClient, err := s.getPDSClient(ctx, session) + if err != nil { + return fmt.Errorf("failed to create PDS client: %w", err) + } + + // Write-forward: delete record from PDS using DPoP-authenticated client + if err := pdsClient.DeleteRecord(ctx, "social.coves.community.block", rkey); err != nil { + if pds.IsAuthError(err) { + return ErrUnauthorized + } return fmt.Errorf("failed to delete block on PDS: %w", err) } diff --git a/tests/e2e/user_signup_test.go b/tests/e2e/user_signup_test.go index 31b2547..dd103b4 100644 --- a/tests/e2e/user_signup_test.go +++ b/tests/e2e/user_signup_test.go @@ -391,15 +391,13 @@ func getProfileViaAPI(did string) (string, string, error) { } var result struct { - DID string `json:"did"` - Profile struct { - Handle string `json:"handle"` - } `json:"profile"` + DID string `json:"did"` + Handle string `json:"handle"` } if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { return "", "", fmt.Errorf("failed to decode response: %w", err) } - return result.DID, result.Profile.Handle, nil + return result.DID, result.Handle, nil } diff --git a/tests/integration/aggregator_e2e_test.go b/tests/integration/aggregator_e2e_test.go index 22ce95a..365ceb8 100644 --- a/tests/integration/aggregator_e2e_test.go +++ b/tests/integration/aggregator_e2e_test.go @@ -68,7 +68,7 @@ func TestAggregator_E2E_WithJetstream(t *testing.T) { identityConfig := identity.DefaultConfig() identityResolver := identity.NewResolver(db, identityConfig) userService := users.NewUserService(userRepo, identityResolver, "http://localhost:3001") - communityService := communities.NewCommunityService(communityRepo, "http://localhost:3001", "did:web:test.coves.social", "coves.social", nil) + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, "http://localhost:3001", "did:web:test.coves.social", "coves.social", nil, nil) aggregatorService := aggregators.NewAggregatorService(aggregatorRepo, communityService) postService := posts.NewPostService(postRepo, communityService, aggregatorService, nil, nil, nil, "http://localhost:3001") diff --git a/tests/integration/author_posts_e2e_test.go b/tests/integration/author_posts_e2e_test.go index 7d00f5f..49b3678 100644 --- a/tests/integration/author_posts_e2e_test.go +++ b/tests/integration/author_posts_e2e_test.go @@ -71,7 +71,7 @@ func TestGetAuthorPosts_E2E_Success(t *testing.T) { // Setup services resolver := identity.NewResolver(db, identity.DefaultConfig()) userService := users.NewUserService(userRepo, resolver, pdsURL) - communityService := communities.NewCommunityService(communityRepo, pdsURL, getTestInstanceDID(), "", nil) + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, pdsURL, getTestInstanceDID(), "", nil, nil) postService := posts.NewPostService(postRepo, communityService, nil, nil, nil, nil, pdsURL) voteService := votes.NewServiceWithPDSFactory(voteRepo, nil, nil, PasswordAuthPDSClientFactory()) @@ -289,7 +289,7 @@ func TestGetAuthorPosts_FilterLogic(t *testing.T) { resolver := identity.NewResolver(db, identity.DefaultConfig()) userService := users.NewUserService(userRepo, resolver, getTestPDSURL()) - communityService := communities.NewCommunityService(communityRepo, getTestPDSURL(), getTestInstanceDID(), "", nil) + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, getTestPDSURL(), getTestInstanceDID(), "", nil, nil) postService := posts.NewPostService(postRepo, communityService, nil, nil, nil, nil, getTestPDSURL()) voteService := votes.NewServiceWithPDSFactory(voteRepo, nil, nil, PasswordAuthPDSClientFactory()) @@ -428,7 +428,7 @@ func TestGetAuthorPosts_ServiceErrors(t *testing.T) { resolver := identity.NewResolver(db, identity.DefaultConfig()) userService := users.NewUserService(userRepo, resolver, getTestPDSURL()) - communityService := communities.NewCommunityService(communityRepo, getTestPDSURL(), getTestInstanceDID(), "", nil) + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, getTestPDSURL(), getTestInstanceDID(), "", nil, nil) postService := posts.NewPostService(postRepo, communityService, nil, nil, nil, nil, getTestPDSURL()) voteService := votes.NewServiceWithPDSFactory(voteRepo, nil, nil, PasswordAuthPDSClientFactory()) @@ -548,7 +548,7 @@ func TestGetAuthorPosts_WithJetstreamIndexing(t *testing.T) { // Setup services resolver := identity.NewResolver(db, identity.DefaultConfig()) userService := users.NewUserService(userRepo, resolver, pdsURL) - communityService := communities.NewCommunityService(communityRepo, pdsURL, getTestInstanceDID(), "", nil) + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, pdsURL, getTestInstanceDID(), "", nil, nil) postService := posts.NewPostService(postRepo, communityService, nil, nil, nil, nil, pdsURL) voteService := votes.NewServiceWithPDSFactory(voteRepo, nil, nil, PasswordAuthPDSClientFactory()) @@ -658,7 +658,7 @@ func TestGetAuthorPosts_CommunityFilter(t *testing.T) { resolver := identity.NewResolver(db, identity.DefaultConfig()) userService := users.NewUserService(userRepo, resolver, getTestPDSURL()) - communityService := communities.NewCommunityService(communityRepo, getTestPDSURL(), getTestInstanceDID(), "", nil) + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, getTestPDSURL(), getTestInstanceDID(), "", nil, nil) postService := posts.NewPostService(postRepo, communityService, nil, nil, nil, nil, getTestPDSURL()) voteService := votes.NewServiceWithPDSFactory(voteRepo, nil, nil, PasswordAuthPDSClientFactory()) diff --git a/tests/integration/block_handle_resolution_test.go b/tests/integration/block_handle_resolution_test.go index 03c84da..237c8ba 100644 --- a/tests/integration/block_handle_resolution_test.go +++ b/tests/integration/block_handle_resolution_test.go @@ -13,8 +13,22 @@ import ( "testing" postgresRepo "Coves/internal/db/postgres" + + "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" ) +// createTestOAuthSessionForBlock creates a mock OAuth session for block handler tests +func createTestOAuthSessionForBlock(did string) *oauth.ClientSessionData { + parsedDID, _ := syntax.ParseDID(did) + return &oauth.ClientSessionData{ + AccountDID: parsedDID, + SessionID: "test-session", + HostURL: "http://localhost:3001", + AccessToken: "test-access-token", + } +} + // TestBlockHandler_HandleResolution tests that the block handler accepts handles // in addition to DIDs and resolves them correctly func TestBlockHandler_HandleResolution(t *testing.T) { @@ -29,12 +43,13 @@ func TestBlockHandler_HandleResolution(t *testing.T) { // Set up repositories and services communityRepo := postgresRepo.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, getTestPDSURL(), getTestInstanceDID(), "coves.social", nil, // No PDS HTTP client for this test + nil, // No PDS factory needed for this test ) blockHandler := community.NewBlockHandler(communityService) @@ -193,9 +208,9 @@ func TestBlockHandler_HandleResolution(t *testing.T) { req := httptest.NewRequest(http.MethodPost, "/xrpc/social.coves.community.blockCommunity", bytes.NewBuffer(reqJSON)) req.Header.Set("Content-Type", "application/json") - // Add auth context so we get past auth checks and test resolution validation - ctx := context.WithValue(req.Context(), middleware.UserDIDKey, "did:plc:test123") - ctx = context.WithValue(ctx, middleware.UserAccessToken, "test-token") + // Add OAuth session context so we get past auth checks and test resolution validation + session := createTestOAuthSessionForBlock("did:plc:test123") + ctx := context.WithValue(req.Context(), middleware.OAuthSessionKey, session) req = req.WithContext(ctx) w := httptest.NewRecorder() @@ -265,12 +280,13 @@ func TestUnblockHandler_HandleResolution(t *testing.T) { // Set up repositories and services communityRepo := postgresRepo.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, getTestPDSURL(), getTestInstanceDID(), "coves.social", nil, + nil, // No PDS factory needed for this test ) blockHandler := community.NewBlockHandler(communityService) diff --git a/tests/integration/community_e2e_test.go b/tests/integration/community_e2e_test.go index 93d9a81..9842d16 100644 --- a/tests/integration/community_e2e_test.go +++ b/tests/integration/community_e2e_test.go @@ -22,6 +22,8 @@ import ( "testing" "time" + oauthlib "github.com/bluesky-social/indigo/atproto/auth/oauth" + "github.com/bluesky-social/indigo/atproto/syntax" "github.com/go-chi/chi/v5" "github.com/gorilla/websocket" _ "github.com/lib/pq" @@ -151,8 +153,8 @@ func TestCommunity_E2E(t *testing.T) { // PDS handles all DID generation and registration automatically provisioner := communities.NewPDSAccountProvisioner(instanceDomain, pdsURL) - // Create service (no longer needs didGen directly - provisioner owns it) - communityService := communities.NewCommunityService(communityRepo, pdsURL, instanceDID, instanceDomain, provisioner) + // Create service with PDS factory for password-based auth in tests + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, pdsURL, instanceDID, instanceDomain, provisioner, CommunityPasswordAuthPDSClientFactory()) if svc, ok := communityService.(interface{ SetPDSAccessToken(string) }); ok { svc.SetPDSAccessToken(accessToken) } @@ -950,7 +952,15 @@ func TestCommunity_E2E(t *testing.T) { t.Logf("Initial subscriber count: %d", initialSubscriberCount) // Subscribe first (using instance access token for instance user, with contentVisibility=3) - subscription, err := communityService.SubscribeToCommunity(ctx, instanceDID, accessToken, community.DID, 3) + // Create a session for the instance user + parsedDID, _ := syntax.ParseDID(instanceDID) + instanceSession := &oauthlib.ClientSessionData{ + AccountDID: parsedDID, + SessionID: "test-session-e2e", + HostURL: pdsURL, + AccessToken: accessToken, + } + subscription, err := communityService.SubscribeToCommunity(ctx, instanceSession, community.DID, 3) if err != nil { t.Fatalf("Failed to subscribe: %v", err) } diff --git a/tests/integration/community_identifier_resolution_test.go b/tests/integration/community_identifier_resolution_test.go index 2db3748..053ce3c 100644 --- a/tests/integration/community_identifier_resolution_test.go +++ b/tests/integration/community_identifier_resolution_test.go @@ -50,12 +50,13 @@ func TestCommunityIdentifierResolution(t *testing.T) { instanceDID = "did:web:" + instanceDomain } - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, instanceDID, instanceDomain, provisioner, + nil, ) // Create a test community to resolve @@ -244,12 +245,13 @@ func TestResolveScopedIdentifier_InputValidation(t *testing.T) { } provisioner := communities.NewPDSAccountProvisioner(instanceDomain, pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, instanceDID, instanceDomain, provisioner, + nil, ) tests := []struct { @@ -421,12 +423,13 @@ func TestIdentifierResolution_ErrorContext(t *testing.T) { } provisioner := communities.NewPDSAccountProvisioner(instanceDomain, pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, instanceDID, instanceDomain, provisioner, + nil, ) t.Run("DID error includes identifier", func(t *testing.T) { @@ -486,12 +489,13 @@ func TestGetCommunity_IdentifierResolution(t *testing.T) { } provisioner := communities.NewPDSAccountProvisioner(instanceDomain, pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, instanceDID, instanceDomain, provisioner, + nil, ) // Create a test community diff --git a/tests/integration/community_provisioning_test.go b/tests/integration/community_provisioning_test.go index e0599d2..a390829 100644 --- a/tests/integration/community_provisioning_test.go +++ b/tests/integration/community_provisioning_test.go @@ -146,12 +146,13 @@ func TestCommunityService_NameValidation(t *testing.T) { repo := postgres.NewCommunityRepository(db) provisioner := communities.NewPDSAccountProvisioner("test.local", "http://localhost:3001") - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, "http://localhost:3001", // pdsURL "did:web:test.local", // instanceDID "test.local", // instanceDomain provisioner, + nil, ) ctx := context.Background() diff --git a/tests/integration/community_service_integration_test.go b/tests/integration/community_service_integration_test.go index 8a4212d..dc5bfaa 100644 --- a/tests/integration/community_service_integration_test.go +++ b/tests/integration/community_service_integration_test.go @@ -57,12 +57,13 @@ func TestCommunityService_CreateWithRealPDS(t *testing.T) { // Create provisioner and service (production code path) // Use coves.social domain (configured in PDS_SERVICE_HANDLE_DOMAINS as c-{name}.coves.social) provisioner := communities.NewPDSAccountProvisioner("coves.social", pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, "did:web:coves.social", "coves.social", provisioner, + nil, ) // Generate unique community name (keep short for DNS label limit) @@ -201,12 +202,13 @@ func TestCommunityService_CreateWithRealPDS(t *testing.T) { t.Run("handles PDS errors gracefully", func(t *testing.T) { provisioner := communities.NewPDSAccountProvisioner("coves.social", pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, "did:web:coves.social", "coves.social", provisioner, + nil, ) // Try to create community with invalid name (should fail validation before PDS) @@ -232,12 +234,13 @@ func TestCommunityService_CreateWithRealPDS(t *testing.T) { t.Run("validates DNS label limits", func(t *testing.T) { provisioner := communities.NewPDSAccountProvisioner("coves.social", pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, "did:web:coves.social", "coves.social", provisioner, + nil, ) // Try 64-char name (exceeds DNS limit of 63) @@ -301,12 +304,13 @@ func TestCommunityService_UpdateWithRealPDS(t *testing.T) { repo := postgres.NewCommunityRepository(db) provisioner := communities.NewPDSAccountProvisioner("coves.social", pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, "did:web:coves.social", "coves.social", provisioner, + nil, ) t.Run("updates community with real PDS", func(t *testing.T) { @@ -492,12 +496,13 @@ func TestPasswordAuthentication(t *testing.T) { repo := postgres.NewCommunityRepository(db) provisioner := communities.NewPDSAccountProvisioner("coves.social", pdsURL) - service := communities.NewCommunityService( + service := communities.NewCommunityServiceWithPDSFactory( repo, pdsURL, "did:web:coves.social", "coves.social", provisioner, + nil, ) t.Run("generated password works for session creation", func(t *testing.T) { diff --git a/tests/integration/community_update_e2e_test.go b/tests/integration/community_update_e2e_test.go index d212896..38047da 100644 --- a/tests/integration/community_update_e2e_test.go +++ b/tests/integration/community_update_e2e_test.go @@ -88,12 +88,13 @@ func TestCommunityUpdateE2E_WithJetstream(t *testing.T) { // Setup services communityRepo := postgres.NewCommunityRepository(db) provisioner := communities.NewPDSAccountProvisioner("coves.social", pdsURL) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, pdsURL, instanceDID, "coves.social", provisioner, + nil, ) consumer := jetstream.NewCommunityEventConsumer(communityRepo, instanceDID, true, identityResolver) diff --git a/tests/integration/feed_test.go b/tests/integration/feed_test.go index 8809dad..edfb098 100644 --- a/tests/integration/feed_test.go +++ b/tests/integration/feed_test.go @@ -29,12 +29,13 @@ func TestGetCommunityFeed_Hot(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -106,12 +107,13 @@ func TestGetCommunityFeed_Top_WithTimeframe(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -182,12 +184,13 @@ func TestGetCommunityFeed_New(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -238,12 +241,13 @@ func TestGetCommunityFeed_Pagination(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -329,12 +333,13 @@ func TestGetCommunityFeed_InvalidCommunity(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -365,12 +370,13 @@ func TestGetCommunityFeed_InvalidCursor(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -421,12 +427,13 @@ func TestGetCommunityFeed_EmptyFeed(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -465,12 +472,13 @@ func TestGetCommunityFeed_LimitValidation(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -518,12 +526,13 @@ func TestGetCommunityFeed_HotPaginationBug(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -619,12 +628,13 @@ func TestGetCommunityFeed_HotCursorPrecision(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -721,12 +731,13 @@ func TestGetCommunityFeed_HotCursorTimeDrift(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) @@ -822,12 +833,13 @@ func TestGetCommunityFeed_BlobURLTransformation(t *testing.T) { // Setup services feedRepo := postgres.NewCommunityFeedRepository(db, "test-cursor-secret") communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) feedService := communityFeeds.NewCommunityFeedService(feedRepo, communityService) handler := communityFeed.NewGetCommunityHandler(feedService, nil, nil) diff --git a/tests/integration/helpers.go b/tests/integration/helpers.go index 3f232d3..796987c 100644 --- a/tests/integration/helpers.go +++ b/tests/integration/helpers.go @@ -4,6 +4,7 @@ import ( "Coves/internal/api/middleware" "Coves/internal/atproto/oauth" "Coves/internal/atproto/pds" + "Coves/internal/core/communities" "Coves/internal/core/users" "Coves/internal/core/votes" "bytes" @@ -443,3 +444,19 @@ func PasswordAuthPDSClientFactory() votes.PDSClientFactory { return pds.NewFromAccessToken(session.HostURL, session.AccountDID.String(), session.AccessToken) } } + +// CommunityPasswordAuthPDSClientFactory creates a PDSClientFactory for communities that uses password-based Bearer auth. +// This is for E2E tests that use createSession instead of OAuth. +// The factory extracts the access token and host URL from the session data. +func CommunityPasswordAuthPDSClientFactory() communities.PDSClientFactory { + return func(ctx context.Context, session *oauthlib.ClientSessionData) (pds.Client, error) { + if session.AccessToken == "" { + return nil, fmt.Errorf("session has no access token") + } + if session.HostURL == "" { + return nil, fmt.Errorf("session has no host URL") + } + + return pds.NewFromAccessToken(session.HostURL, session.AccountDID.String(), session.AccessToken) + } +} diff --git a/tests/integration/post_creation_test.go b/tests/integration/post_creation_test.go index 71f2215..152eaf9 100644 --- a/tests/integration/post_creation_test.go +++ b/tests/integration/post_creation_test.go @@ -35,12 +35,13 @@ func TestPostCreation_Basic(t *testing.T) { communityRepo := postgres.NewCommunityRepository(db) // Note: Provisioner not needed for this test (we're not actually creating communities) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, // provisioner + nil, // pdsClientFactory ) postRepo := postgres.NewPostRepository(db) diff --git a/tests/integration/post_e2e_test.go b/tests/integration/post_e2e_test.go index 0e7848a..fdaa8ad 100644 --- a/tests/integration/post_e2e_test.go +++ b/tests/integration/post_e2e_test.go @@ -394,12 +394,13 @@ func TestPostCreation_E2E_LivePDS(t *testing.T) { provisioner := communities.NewPDSAccountProvisioner(instanceDomain, pdsURL) // Setup community service with real PDS provisioner - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, pdsURL, instanceDID, instanceDomain, provisioner, // ✅ Real provisioner for creating communities on PDS + nil, // No PDS factory needed - no subscribe/block in this test ) postService := posts.NewPostService(postRepo, communityService, nil, nil, nil, nil, pdsURL) // nil aggregatorService, blobService, unfurlService, blueskyService for user-only tests diff --git a/tests/integration/post_handler_test.go b/tests/integration/post_handler_test.go index 5ff06bc..709aa1d 100644 --- a/tests/integration/post_handler_test.go +++ b/tests/integration/post_handler_test.go @@ -32,12 +32,13 @@ func TestPostHandler_SecurityValidation(t *testing.T) { // Setup services communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) postRepo := postgres.NewPostRepository(db) @@ -400,12 +401,13 @@ func TestPostHandler_SpecialCharacters(t *testing.T) { // Setup services communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) postRepo := postgres.NewPostRepository(db) @@ -484,12 +486,13 @@ func TestPostService_DIDValidationSecurity(t *testing.T) { // Setup services communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) postRepo := postgres.NewPostRepository(db) diff --git a/tests/integration/post_thumb_validation_test.go b/tests/integration/post_thumb_validation_test.go index 0e70c35..fd0f62d 100644 --- a/tests/integration/post_thumb_validation_test.go +++ b/tests/integration/post_thumb_validation_test.go @@ -55,12 +55,13 @@ func TestPostHandler_ThumbValidation(t *testing.T) { // Setup services communityRepo := postgres.NewCommunityRepository(db) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) postRepo := postgres.NewPostRepository(db) diff --git a/tests/integration/post_unfurl_test.go b/tests/integration/post_unfurl_test.go index 8eb10e0..bb5f3a0 100644 --- a/tests/integration/post_unfurl_test.go +++ b/tests/integration/post_unfurl_test.go @@ -51,12 +51,13 @@ func TestPostUnfurl_Streamable(t *testing.T) { unfurl.WithCacheTTL(24*time.Hour), ) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) postService := posts.NewPostService( @@ -348,12 +349,13 @@ func TestPostUnfurl_UnsupportedURL(t *testing.T) { identityResolver := identity.NewResolver(db, identityConfig) userService := users.NewUserService(userRepo, identityResolver, "http://localhost:3001") - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) // Create post service WITHOUT unfurl service @@ -456,12 +458,13 @@ func TestPostUnfurl_UserProvidedMetadata(t *testing.T) { unfurl.WithCacheTTL(24*time.Hour), ) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) postService := posts.NewPostService( @@ -568,12 +571,13 @@ func TestPostUnfurl_MissingEmbedType(t *testing.T) { unfurl.WithTimeout(30*time.Second), ) - communityService := communities.NewCommunityService( + communityService := communities.NewCommunityServiceWithPDSFactory( communityRepo, "http://localhost:3001", "did:web:test.coves.social", "test.coves.social", nil, + nil, ) postService := posts.NewPostService( diff --git a/tests/integration/user_journey_e2e_test.go b/tests/integration/user_journey_e2e_test.go index 0d7cf29..9f55e1e 100644 --- a/tests/integration/user_journey_e2e_test.go +++ b/tests/integration/user_journey_e2e_test.go @@ -128,7 +128,7 @@ func TestFullUserJourney_E2E(t *testing.T) { } provisioner := communities.NewPDSAccountProvisioner(instanceDomain, pdsURL) - communityService := communities.NewCommunityService(communityRepo, pdsURL, instanceDID, instanceDomain, provisioner) + communityService := communities.NewCommunityServiceWithPDSFactory(communityRepo, pdsURL, instanceDID, instanceDomain, provisioner, CommunityPasswordAuthPDSClientFactory()) postService := posts.NewPostService(postRepo, communityService, nil, nil, nil, nil, pdsURL) timelineService := timelineCore.NewTimelineService(timelineRepo)