From 168313c757583c753ca7d4678cab3a46804cbdfa Mon Sep 17 00:00:00 2001 From: Max Date: Tue, 14 Oct 2025 14:15:41 +0800 Subject: [PATCH] Update AcceptInvitation method to support optional user ID parameter - Modified the AcceptInvitation method to accept an optional user ID parameter, allowing for updates to the user ID when accepting invitations without an existing user ID. - Adjusted related tests to include the new user ID parameter, ensuring comprehensive coverage of invitation acceptance scenarios. - Updated the team invitation acceptance logic to utilize the new parameter, enhancing the invitation flow and user management capabilities. --- openapi/oauth/providers/user/member.go | 8 +++++++- openapi/oauth/providers/user/member_test.go | 8 ++++---- openapi/oauth/providers/user/team_test.go | 2 +- openapi/oauth/types/interfaces.go | 2 +- openapi/user/team_invitation.go | 21 ++------------------- 5 files changed, 15 insertions(+), 26 deletions(-) diff --git a/openapi/oauth/providers/user/member.go b/openapi/oauth/providers/user/member.go index b2b6bb67..097981f4 100644 --- a/openapi/oauth/providers/user/member.go +++ b/openapi/oauth/providers/user/member.go @@ -238,7 +238,8 @@ func (u *DefaultUser) AddMember(ctx context.Context, teamID string, userID strin } // AcceptInvitation accepts a team invitation -func (u *DefaultUser) AcceptInvitation(ctx context.Context, invitationID string, invitationToken string) error { +// userID can be empty - if provided and invitation doesn't have user_id, it will be updated +func (u *DefaultUser) AcceptInvitation(ctx context.Context, invitationID string, invitationToken string, userID string) error { // Find member by invitation_id and token m := model.Select(u.memberModel) members, err := m.Get(model.QueryParam{ @@ -277,6 +278,11 @@ func (u *DefaultUser) AcceptInvitation(ctx context.Context, invitationID string, "invitation_token": nil, // Clear the token } + // If invitation doesn't have a user_id (unregistered user invitation), update it with provided userID + if userID != "" && (member["user_id"] == nil || member["user_id"] == "") { + updateData["user_id"] = userID + } + affected, err := m.UpdateWhere(model.QueryParam{ Wheres: []model.QueryWhere{ {Column: "id", Value: memberID}, diff --git a/openapi/oauth/providers/user/member_test.go b/openapi/oauth/providers/user/member_test.go index f6937e24..8257be9b 100644 --- a/openapi/oauth/providers/user/member_test.go +++ b/openapi/oauth/providers/user/member_test.go @@ -278,7 +278,7 @@ func TestMemberInvitationFlow(t *testing.T) { // Test AcceptInvitation t.Run("AcceptInvitation", func(t *testing.T) { - err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken) + err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken, "") assert.NoError(t, err) // Verify member status changed to active @@ -295,14 +295,14 @@ func TestMemberInvitationFlow(t *testing.T) { // Test AcceptInvitation with invalid token t.Run("AcceptInvitation_InvalidToken", func(t *testing.T) { - err := testProvider.AcceptInvitation(ctx, invitationID, "invalid-token") + err := testProvider.AcceptInvitation(ctx, invitationID, "invalid-token", "") assert.Error(t, err) assert.Contains(t, err.Error(), "invitation not found") }) // Test AcceptInvitation with already accepted token t.Run("AcceptInvitation_AlreadyAccepted", func(t *testing.T) { - err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken) + err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken, "") assert.Error(t, err) assert.Contains(t, err.Error(), "invitation not found") }) @@ -750,7 +750,7 @@ func TestMemberInvitationExpiry(t *testing.T) { // Test AcceptInvitation with expired token t.Run("AcceptInvitation_ExpiredToken", func(t *testing.T) { - err := testProvider.AcceptInvitation(ctx, invitationID, "expired-token-"+testUUID) + err := testProvider.AcceptInvitation(ctx, invitationID, "expired-token-"+testUUID, "") assert.Error(t, err) assert.Contains(t, err.Error(), "invitation has expired") }) diff --git a/openapi/oauth/providers/user/team_test.go b/openapi/oauth/providers/user/team_test.go index ee2cf327..433c757a 100644 --- a/openapi/oauth/providers/user/team_test.go +++ b/openapi/oauth/providers/user/team_test.go @@ -381,7 +381,7 @@ func TestTeamMemberOperations(t *testing.T) { assert.NotEmpty(t, invitationID) // Accept the invitation - err = testProvider.AcceptInvitation(ctx, invitationID, invitationToken) + err = testProvider.AcceptInvitation(ctx, invitationID, invitationToken, "") assert.NoError(t, err) // Verify member status changed to active diff --git a/openapi/oauth/types/interfaces.go b/openapi/oauth/types/interfaces.go index eb865498..39b2dcce 100644 --- a/openapi/oauth/types/interfaces.go +++ b/openapi/oauth/types/interfaces.go @@ -308,7 +308,7 @@ type UserProvider interface { // Member Invitation Management AddMember(ctx context.Context, teamID string, userID string, roleID string, invitedBy string) (int64, error) - AcceptInvitation(ctx context.Context, invitationID string, invitationToken string) error + AcceptInvitation(ctx context.Context, invitationID string, invitationToken string, userID string) error // Robot Member Operations CreateRobotMember(ctx context.Context, teamID string, robotData maps.MapStrAny) (int64, error) diff --git a/openapi/user/team_invitation.go b/openapi/user/team_invitation.go index 2f7d21b1..d9a9fe4a 100644 --- a/openapi/user/team_invitation.go +++ b/openapi/user/team_invitation.go @@ -528,25 +528,8 @@ func GinTeamInvitationAccept(c *gin.Context) { return } - // If invitation doesn't have a user_id (unregistered user invitation), update it with current user - if invitationData["user_id"] == nil || invitationData["user_id"] == "" { - updateData := maps.MapStrAny{ - "user_id": userID, - } - err = provider.UpdateMemberByInvitationID(ctx, invitationID, updateData) - if err != nil { - log.Error("Failed to update invitation with user_id: %v", err) - errorResp := &response.ErrorResponse{ - Code: response.ErrServerError.Code, - ErrorDescription: "Failed to process invitation", - } - response.RespondWithError(c, response.StatusInternalServerError, errorResp) - return - } - } - - // Accept the invitation - err = provider.AcceptInvitation(ctx, invitationID, req.Token) + // Accept the invitation (will update user_id if invitation doesn't have one) + err = provider.AcceptInvitation(ctx, invitationID, req.Token, userID) if err != nil { log.Error("Failed to accept invitation: %v", err) // Check error type for appropriate response