Merge pull request #1199 from trheyi/main

Update AcceptInvitation method to support optional user ID parameter
This commit is contained in:
Max 2025-10-14 14:16:07 +08:00 committed by GitHub
commit 8ba7730fb1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 15 additions and 26 deletions

View file

@ -238,7 +238,8 @@ func (u *DefaultUser) AddMember(ctx context.Context, teamID string, userID strin
} }
// AcceptInvitation accepts a team invitation // 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 // Find member by invitation_id and token
m := model.Select(u.memberModel) m := model.Select(u.memberModel)
members, err := m.Get(model.QueryParam{ 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 "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{ affected, err := m.UpdateWhere(model.QueryParam{
Wheres: []model.QueryWhere{ Wheres: []model.QueryWhere{
{Column: "id", Value: memberID}, {Column: "id", Value: memberID},

View file

@ -278,7 +278,7 @@ func TestMemberInvitationFlow(t *testing.T) {
// Test AcceptInvitation // Test AcceptInvitation
t.Run("AcceptInvitation", func(t *testing.T) { t.Run("AcceptInvitation", func(t *testing.T) {
err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken) err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken, "")
assert.NoError(t, err) assert.NoError(t, err)
// Verify member status changed to active // Verify member status changed to active
@ -295,14 +295,14 @@ func TestMemberInvitationFlow(t *testing.T) {
// Test AcceptInvitation with invalid token // Test AcceptInvitation with invalid token
t.Run("AcceptInvitation_InvalidToken", func(t *testing.T) { 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.Error(t, err)
assert.Contains(t, err.Error(), "invitation not found") assert.Contains(t, err.Error(), "invitation not found")
}) })
// Test AcceptInvitation with already accepted token // Test AcceptInvitation with already accepted token
t.Run("AcceptInvitation_AlreadyAccepted", func(t *testing.T) { 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.Error(t, err)
assert.Contains(t, err.Error(), "invitation not found") assert.Contains(t, err.Error(), "invitation not found")
}) })
@ -750,7 +750,7 @@ func TestMemberInvitationExpiry(t *testing.T) {
// Test AcceptInvitation with expired token // Test AcceptInvitation with expired token
t.Run("AcceptInvitation_ExpiredToken", func(t *testing.T) { 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.Error(t, err)
assert.Contains(t, err.Error(), "invitation has expired") assert.Contains(t, err.Error(), "invitation has expired")
}) })

View file

@ -381,7 +381,7 @@ func TestTeamMemberOperations(t *testing.T) {
assert.NotEmpty(t, invitationID) assert.NotEmpty(t, invitationID)
// Accept the invitation // Accept the invitation
err = testProvider.AcceptInvitation(ctx, invitationID, invitationToken) err = testProvider.AcceptInvitation(ctx, invitationID, invitationToken, "")
assert.NoError(t, err) assert.NoError(t, err)
// Verify member status changed to active // Verify member status changed to active

View file

@ -308,7 +308,7 @@ type UserProvider interface {
// Member Invitation Management // Member Invitation Management
AddMember(ctx context.Context, teamID string, userID string, roleID string, invitedBy string) (int64, error) 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 // Robot Member Operations
CreateRobotMember(ctx context.Context, teamID string, robotData maps.MapStrAny) (int64, error) CreateRobotMember(ctx context.Context, teamID string, robotData maps.MapStrAny) (int64, error)

View file

@ -528,25 +528,8 @@ func GinTeamInvitationAccept(c *gin.Context) {
return return
} }
// If invitation doesn't have a user_id (unregistered user invitation), update it with current user // Accept the invitation (will update user_id if invitation doesn't have one)
if invitationData["user_id"] == nil || invitationData["user_id"] == "" { err = provider.AcceptInvitation(ctx, invitationID, req.Token, userID)
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)
if err != nil { if err != nil {
log.Error("Failed to accept invitation: %v", err) log.Error("Failed to accept invitation: %v", err)
// Check error type for appropriate response // Check error type for appropriate response