From 721f47c3458ede33f4600fafe1663dbf02e83004 Mon Sep 17 00:00:00 2001 From: Max Date: Tue, 14 Oct 2025 11:58:20 +0800 Subject: [PATCH] Enhance invitation acceptance flow with invitation ID support - Updated the AcceptInvitation method to require both invitation ID and token, improving the invitation acceptance process. - Modified related tests to accommodate the new invitation ID parameter, ensuring comprehensive coverage of acceptance scenarios. - Enhanced the invitation acceptance endpoint to validate invitation ID, providing clearer error handling for invalid or expired invitations. - Refactored tests to include detailed scenarios for accepting invitations, including success and failure cases, ensuring robust testing of the invitation flow. --- openapi/oauth/providers/user/member.go | 5 +- openapi/oauth/providers/user/member_test.go | 21 +- openapi/oauth/providers/user/team_test.go | 6 +- openapi/oauth/types/interfaces.go | 2 +- openapi/tests/testutils/testutils.go | 101 ++++++ openapi/tests/user/invitation_test.go | 361 ++++++++++++++++++++ openapi/tests/user/team_test.go | 38 +-- openapi/user/team_invitation.go | 153 +++++++++ openapi/user/user.go | 3 +- 9 files changed, 652 insertions(+), 38 deletions(-) diff --git a/openapi/oauth/providers/user/member.go b/openapi/oauth/providers/user/member.go index c1b00fb2..b2b6bb67 100644 --- a/openapi/oauth/providers/user/member.go +++ b/openapi/oauth/providers/user/member.go @@ -238,12 +238,13 @@ func (u *DefaultUser) AddMember(ctx context.Context, teamID string, userID strin } // AcceptInvitation accepts a team invitation -func (u *DefaultUser) AcceptInvitation(ctx context.Context, invitationToken string) error { - // Find member by invitation token +func (u *DefaultUser) AcceptInvitation(ctx context.Context, invitationID string, invitationToken string) error { + // Find member by invitation_id and token m := model.Select(u.memberModel) members, err := m.Get(model.QueryParam{ Select: []interface{}{"id", "team_id", "user_id", "status", "invitation_expires_at"}, Wheres: []model.QueryWhere{ + {Column: "invitation_id", Value: invitationID}, {Column: "invitation_token", Value: invitationToken}, {Column: "status", Value: "pending"}, }, diff --git a/openapi/oauth/providers/user/member_test.go b/openapi/oauth/providers/user/member_test.go index b6e129d1..f6937e24 100644 --- a/openapi/oauth/providers/user/member_test.go +++ b/openapi/oauth/providers/user/member_test.go @@ -243,6 +243,7 @@ func TestMemberInvitationFlow(t *testing.T) { assert.NoError(t, err) var invitationToken string + var invitationID string // Test AddMember (invitation-based) t.Run("AddMember", func(t *testing.T) { @@ -256,11 +257,13 @@ func TestMemberInvitationFlow(t *testing.T) { assert.Equal(t, "pending", member["status"]) assert.Equal(t, ownerUser, member["invited_by"]) - // Get invitation token for acceptance test + // Get invitation token and invitation_id for acceptance test memberDetail, err := testProvider.GetMemberDetail(ctx, teamID, inviteeUser) assert.NoError(t, err) invitationToken = memberDetail["invitation_token"].(string) assert.NotEmpty(t, invitationToken) + invitationID = memberDetail["invitation_id"].(string) + assert.NotEmpty(t, invitationID) // Verify invitation expiry is set assert.NotNil(t, memberDetail["invitation_expires_at"]) @@ -275,7 +278,7 @@ func TestMemberInvitationFlow(t *testing.T) { // Test AcceptInvitation t.Run("AcceptInvitation", func(t *testing.T) { - err := testProvider.AcceptInvitation(ctx, invitationToken) + err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken) assert.NoError(t, err) // Verify member status changed to active @@ -292,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, "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, invitationToken) + err := testProvider.AcceptInvitation(ctx, invitationID, invitationToken) assert.Error(t, err) assert.Contains(t, err.Error(), "invitation not found") }) @@ -736,12 +739,18 @@ func TestMemberInvitationExpiry(t *testing.T) { "invitation_expires_at": expiredTime, // Expired 2 hours ago } - _, err = testProvider.CreateMember(ctx, memberData) + memberID, err := testProvider.CreateMember(ctx, memberData) assert.NoError(t, err) + // Get the invitation_id + member, err := testProvider.GetMemberByID(ctx, memberID) + assert.NoError(t, err) + invitationID := member["invitation_id"].(string) + assert.NotEmpty(t, invitationID) + // Test AcceptInvitation with expired token t.Run("AcceptInvitation_ExpiredToken", func(t *testing.T) { - err := testProvider.AcceptInvitation(ctx, "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 74b1633d..ee2cf327 100644 --- a/openapi/oauth/providers/user/team_test.go +++ b/openapi/oauth/providers/user/team_test.go @@ -372,14 +372,16 @@ func TestTeamMemberOperations(t *testing.T) { // Test AcceptInvitation t.Run("AcceptInvitation", func(t *testing.T) { - // First get the invitation token + // First get the invitation token and invitation_id member, err := testProvider.GetMemberDetail(ctx, teamID, memberUser) assert.NoError(t, err) invitationToken := member["invitation_token"].(string) assert.NotEmpty(t, invitationToken) + invitationID := member["invitation_id"].(string) + assert.NotEmpty(t, invitationID) // Accept the invitation - err = testProvider.AcceptInvitation(ctx, 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 e4ed297c..eb865498 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, invitationToken string) error + AcceptInvitation(ctx context.Context, invitationID string, invitationToken string) error // Robot Member Operations CreateRobotMember(ctx context.Context, teamID string, robotData maps.MapStrAny) (int64, error) diff --git a/openapi/tests/testutils/testutils.go b/openapi/tests/testutils/testutils.go index 62383de5..a8f2087d 100644 --- a/openapi/tests/testutils/testutils.go +++ b/openapi/tests/testutils/testutils.go @@ -701,6 +701,107 @@ func generateCodeChallenge(codeVerifier string) string { return base64.URLEncoding.WithPadding(base64.NoPadding).EncodeToString(hash[:]) } +// ObtainTokenForUser creates a token for a specific user ID. +// This function creates the OAuth fingerprint mapping and issues tokens for the given userID. +// +// AI ASSISTANT INSTRUCTIONS: +// Use this function when you need to issue a token for a pre-existing user ID. +// This is useful for testing scenarios where user records already exist in the database. +// +// Usage pattern: +// +// func TestWithExistingUser(t *testing.T) { +// serverURL := Prepare(t) +// defer Clean() +// +// // Create user in database first +// userID := "user_123" +// // ... create user record in DB ... +// +// // Register a test client +// client := RegisterTestClient(t, "Test Client", []string{"https://localhost/callback"}) +// defer CleanupTestClient(t, client.ClientID) +// +// // Obtain token for the existing user +// tokenInfo := ObtainTokenForUser(t, client.ClientID, client.ClientSecret, userID, "openid profile") +// +// // Now use tokenInfo.AccessToken to make authenticated requests +// } +// +// PARAMETERS: +// - t: The test instance for error reporting +// - clientID: The OAuth client ID (from RegisterTestClient) +// - clientSecret: The OAuth client secret (from RegisterTestClient) +// - userID: The user ID to issue the token for (must exist in database) +// - scope: The requested OAuth scope (e.g., "openid profile email") +// +// RETURN VALUE: +// Returns TokenInfo struct containing: +// - AccessToken: The access token for API calls +// - RefreshToken: The refresh token for token renewal +// - TokenType: The token type (usually "Bearer") +// - ExpiresIn: Token expiration time in seconds +// - Scope: The granted scope +// - ClientID: The client ID used to obtain the token +// - UserID: The user ID the token was issued for +// +// WHAT THIS FUNCTION DOES: +// 1. Creates a fingerprint mapping: clientID + subject -> userID +// 2. Issues access and refresh tokens for the subject +// 3. Returns all token information needed for authenticated API testing +// +// ERROR HANDLING: +// If token creation fails, the test will fail immediately with a descriptive error message. +func ObtainTokenForUser(t *testing.T, clientID, clientSecret, userID, scope string) *TokenInfo { + testMutex.RLock() + server := openapi.Server + testMutex.RUnlock() + + if server == nil || server.OAuth == nil { + t.Fatal("OpenAPI server not initialized. Call Prepare(t) first.") + } + + // Access the global OAuth service + oauthService := oauth.OAuth + if oauthService == nil { + t.Fatal("Global OAuth service not initialized") + } + + // Create subject (fingerprint) for this user + // This sets up the fingerprint mapping: clientID:subject -> userID + subject, err := oauthService.Subject(clientID, userID) + if err != nil { + t.Fatalf("Failed to create user subject: %v", err) + } + + t.Logf("Created fingerprint mapping: clientID=%s, userID=%s, subject=%s", clientID, userID, subject) + + // Create access token + accessToken, err := oauthService.MakeAccessToken(clientID, scope, subject, 3600) + if err != nil { + t.Fatalf("Failed to create access token: %v", err) + } + + // Create refresh token + refreshToken, err := oauthService.MakeRefreshToken(clientID, scope, subject, 7200) + if err != nil { + t.Fatalf("Failed to create refresh token: %v", err) + } + + tokenInfo := &TokenInfo{ + AccessToken: accessToken, + RefreshToken: refreshToken, + TokenType: "Bearer", + ExpiresIn: 3600, + Scope: scope, + ClientID: clientID, + UserID: userID, + } + + t.Logf("Issued token for user %s (subject: %s)", userID, subject) + return tokenInfo +} + // createTestUser creates a test user and sets up proper fingerprint mapping for OAuth authentication func createTestUser(t *testing.T, server *openapi.OpenAPI, clientID string) (string, string) { if server.OAuth == nil { diff --git a/openapi/tests/user/invitation_test.go b/openapi/tests/user/invitation_test.go index fc65441d..392375ad 100644 --- a/openapi/tests/user/invitation_test.go +++ b/openapi/tests/user/invitation_test.go @@ -2,6 +2,7 @@ package user_test import ( "bytes" + "context" "encoding/json" "fmt" "io" @@ -13,6 +14,7 @@ import ( "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/yaoapp/yao/openapi" + "github.com/yaoapp/yao/openapi/oauth" "github.com/yaoapp/yao/openapi/tests/testutils" "github.com/yaoapp/yao/openapi/user" ) @@ -839,8 +841,367 @@ func TestInvitationDelete(t *testing.T) { }) } +// TestInvitationAccept tests the POST /user/teams/invitations/:invitation_id/accept endpoint +func TestInvitationAccept(t *testing.T) { + serverURL := testutils.Prepare(t) + defer testutils.Clean() + + // Get base URL from server config + baseURL := "" + if openapi.Server != nil && openapi.Server.Config != nil { + baseURL = openapi.Server.Config.BaseURL + } + + // Use UUID to ensure unique identifiers + testUUID := strings.ReplaceAll(uuid.New().String(), "-", "")[:8] + + // Step 1: Create two users A and B in database + t.Logf("Step 1: Create users A and B") + userA := fmt.Sprintf("user_a_%s", testUUID) + userB := fmt.Sprintf("user_b_%s", testUUID) + + // Create users and get actual user IDs returned by provider + actualUserA := createUserInDB(t, userA) + actualUserB := createUserInDB(t, userB) + + // Use the actual user IDs returned by CreateUser + userA = actualUserA + userB = actualUserB + t.Logf(" - Created users: A=%s, B=%s", userA, userB) + + // Step 2: Issue token for user A and create team + t.Logf("Step 2: Issue token for user A and create team") + clientA := testutils.RegisterTestClient(t, "User A Client "+testUUID, []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, clientA.ClientID) + + tokenA := testutils.ObtainTokenForUser(t, clientA.ClientID, clientA.ClientSecret, userA, "openid profile") + teamID, invitationID := setupTeamAndInvitation(t, serverURL, baseURL, tokenA.AccessToken, userA, userB, testUUID) + t.Logf(" - Team created with ID: %s", teamID) + t.Logf(" - Invitation created with ID: %s", invitationID) + + // Step 3: Issue token for user B + t.Logf("Step 3: Issue token for user B") + clientB := testutils.RegisterTestClient(t, "User B Client "+testUUID, []string{"https://localhost/callback"}) + defer testutils.CleanupTestClient(t, clientB.ClientID) + + tokenB := testutils.ObtainTokenForUser(t, clientB.ClientID, clientB.ClientSecret, userB, "openid profile") + + // Test successful accept invitation + t.Run("AcceptInvitation_Success", func(t *testing.T) { + // Get invitation details to retrieve token + getURL := fmt.Sprintf("%s%s/user/teams/%s/invitations/%s", serverURL, baseURL, teamID, invitationID) + getReq, err := http.NewRequest("GET", getURL, nil) + assert.NoError(t, err) + getReq.Header.Set("Authorization", "Bearer "+tokenA.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + getResp, err := client.Do(getReq) + assert.NoError(t, err) + defer getResp.Body.Close() + + assert.Equal(t, http.StatusOK, getResp.StatusCode) + + var invitation user.InvitationDetailResponse + err = json.NewDecoder(getResp.Body).Decode(&invitation) + assert.NoError(t, err) + + invitationToken := invitation.InvitationToken + assert.NotEmpty(t, invitationToken) + + // User B accepts the invitation + t.Logf(" - User B accepting invitation with token") + acceptData := map[string]interface{}{ + "token": invitationToken, + } + jsonData, _ := json.Marshal(acceptData) + + acceptURL := fmt.Sprintf("%s%s/user/teams/invitations/%s/accept", serverURL, baseURL, invitationID) + acceptReq, err := http.NewRequest("POST", acceptURL, bytes.NewBuffer(jsonData)) + assert.NoError(t, err) + acceptReq.Header.Set("Content-Type", "application/json") + acceptReq.Header.Set("Authorization", "Bearer "+tokenB.AccessToken) + + acceptResp, err := client.Do(acceptReq) + assert.NoError(t, err) + defer acceptResp.Body.Close() + + // Read response body first for better error message + var result map[string]interface{} + err = json.NewDecoder(acceptResp.Body).Decode(&result) + assert.NoError(t, err) + + // Check status code and provide helpful error message if failed + if acceptResp.StatusCode != http.StatusOK { + t.Fatalf("Accept invitation failed: status=%d, body=%v", acceptResp.StatusCode, result) + } + + // Check for standard LoginResponse fields + assert.Contains(t, result, "access_token") + assert.Contains(t, result, "refresh_token") + assert.Contains(t, result, "token_type") + assert.Contains(t, result, "expires_in") + assert.Contains(t, result, "user_id") + assert.Contains(t, result, "id_token") + + // Verify user_id matches invitee + assert.Equal(t, userB, result["user_id"]) + + // Verify tokens are valid (non-empty) + assert.NotEmpty(t, result["access_token"]) + assert.NotEmpty(t, result["refresh_token"]) + assert.Equal(t, "Bearer", result["token_type"]) + assert.Greater(t, int(result["expires_in"].(float64)), 0) + }) + + // Test accept invitation with invalid token + t.Run("AcceptInvitation_InvalidToken", func(t *testing.T) { + // Create new invitation for this test + _, invID := setupTeamAndInvitation(t, serverURL, baseURL, tokenA.AccessToken, userA, userB, testUUID+"_inv") + + // Try to accept with invalid token + acceptData := map[string]interface{}{ + "token": "invalid-token-12345", + } + jsonData, _ := json.Marshal(acceptData) + + acceptURL := fmt.Sprintf("%s%s/user/teams/invitations/%s/accept", serverURL, baseURL, invID) + acceptReq, err := http.NewRequest("POST", acceptURL, bytes.NewBuffer(jsonData)) + assert.NoError(t, err) + acceptReq.Header.Set("Content-Type", "application/json") + acceptReq.Header.Set("Authorization", "Bearer "+tokenB.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + acceptResp, err := client.Do(acceptReq) + assert.NoError(t, err) + defer acceptResp.Body.Close() + + assert.Equal(t, http.StatusNotFound, acceptResp.StatusCode) + }) + + // Test accept invitation with non-existent invitation_id + t.Run("AcceptInvitation_NonExistentInvitation", func(t *testing.T) { + acceptData := map[string]interface{}{ + "token": "some-token", + } + jsonData, _ := json.Marshal(acceptData) + + acceptURL := fmt.Sprintf("%s%s/user/teams/invitations/non-existent-inv/accept", serverURL, baseURL) + acceptReq, err := http.NewRequest("POST", acceptURL, bytes.NewBuffer(jsonData)) + assert.NoError(t, err) + acceptReq.Header.Set("Content-Type", "application/json") + acceptReq.Header.Set("Authorization", "Bearer "+tokenB.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + acceptResp, err := client.Do(acceptReq) + assert.NoError(t, err) + defer acceptResp.Body.Close() + + assert.Equal(t, http.StatusNotFound, acceptResp.StatusCode) + }) + + // Test accept invitation without authentication + t.Run("AcceptInvitation_Unauthorized", func(t *testing.T) { + // Create new invitation for this test + _, invID := setupTeamAndInvitation(t, serverURL, baseURL, tokenA.AccessToken, userA, userB, testUUID+"_unauth") + + acceptData := map[string]interface{}{ + "token": "some-token", + } + jsonData, _ := json.Marshal(acceptData) + + acceptURL := fmt.Sprintf("%s%s/user/teams/invitations/%s/accept", serverURL, baseURL, invID) + acceptReq, err := http.NewRequest("POST", acceptURL, bytes.NewBuffer(jsonData)) + assert.NoError(t, err) + acceptReq.Header.Set("Content-Type", "application/json") + // No Authorization header + + client := &http.Client{Timeout: 10 * time.Second} + acceptResp, err := client.Do(acceptReq) + assert.NoError(t, err) + defer acceptResp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, acceptResp.StatusCode) + }) + + // Test accept invitation without token in request body + t.Run("AcceptInvitation_MissingToken", func(t *testing.T) { + // Create new invitation for this test + _, invID := setupTeamAndInvitation(t, serverURL, baseURL, tokenA.AccessToken, userA, userB, testUUID+"_missing") + + acceptData := map[string]interface{}{ + // Missing token field + } + jsonData, _ := json.Marshal(acceptData) + + acceptURL := fmt.Sprintf("%s%s/user/teams/invitations/%s/accept", serverURL, baseURL, invID) + acceptReq, err := http.NewRequest("POST", acceptURL, bytes.NewBuffer(jsonData)) + assert.NoError(t, err) + acceptReq.Header.Set("Content-Type", "application/json") + acceptReq.Header.Set("Authorization", "Bearer "+tokenB.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + acceptResp, err := client.Do(acceptReq) + assert.NoError(t, err) + defer acceptResp.Body.Close() + + assert.Equal(t, http.StatusBadRequest, acceptResp.StatusCode) + }) + + // Test accept already accepted invitation + t.Run("AcceptInvitation_AlreadyAccepted", func(t *testing.T) { + // Create new invitation for this test + tID, invID := setupTeamAndInvitation(t, serverURL, baseURL, tokenA.AccessToken, userA, userB, testUUID+"_accepted") + + // Get invitation token + getURL := fmt.Sprintf("%s%s/user/teams/%s/invitations/%s", serverURL, baseURL, tID, invID) + getReq, err := http.NewRequest("GET", getURL, nil) + assert.NoError(t, err) + getReq.Header.Set("Authorization", "Bearer "+tokenA.AccessToken) + + client := &http.Client{Timeout: 10 * time.Second} + getResp, err := client.Do(getReq) + assert.NoError(t, err) + defer getResp.Body.Close() + + var invitation user.InvitationDetailResponse + json.NewDecoder(getResp.Body).Decode(&invitation) + invitationToken := invitation.InvitationToken + + // Accept the invitation first time + acceptData := map[string]interface{}{ + "token": invitationToken, + } + jsonData, _ := json.Marshal(acceptData) + + acceptURL := fmt.Sprintf("%s%s/user/teams/invitations/%s/accept", serverURL, baseURL, invID) + acceptReq, err := http.NewRequest("POST", acceptURL, bytes.NewBuffer(jsonData)) + assert.NoError(t, err) + acceptReq.Header.Set("Content-Type", "application/json") + acceptReq.Header.Set("Authorization", "Bearer "+tokenB.AccessToken) + + acceptResp, err := client.Do(acceptReq) + assert.NoError(t, err) + acceptResp.Body.Close() + + assert.Equal(t, http.StatusOK, acceptResp.StatusCode) + + // Try to accept again (should fail) + acceptReq2, err := http.NewRequest("POST", acceptURL, bytes.NewBuffer(jsonData)) + assert.NoError(t, err) + acceptReq2.Header.Set("Content-Type", "application/json") + acceptReq2.Header.Set("Authorization", "Bearer "+tokenB.AccessToken) + + acceptResp2, err := client.Do(acceptReq2) + assert.NoError(t, err) + defer acceptResp2.Body.Close() + + assert.Equal(t, http.StatusNotFound, acceptResp2.StatusCode) + }) +} + // Helper functions +// setupTeamAndInvitation creates a team and invitation for testing by calling HTTP APIs +// This simulates the complete flow including OAuth Guard middleware +// Returns teamID and invitationID +func setupTeamAndInvitation(t *testing.T, serverURL, baseURL, accessToken, ownerUserID, inviteeUserID, testUUID string) (string, string) { + client := &http.Client{Timeout: 10 * time.Second} + + // Step 1: Create team via HTTP API + teamName := fmt.Sprintf("Team_%s", testUUID) + teamData := map[string]interface{}{ + "name": teamName, + "description": "Test team for invitation acceptance", + } + teamJSON, _ := json.Marshal(teamData) + + createTeamURL := fmt.Sprintf("%s%s/user/teams", serverURL, baseURL) + req, err := http.NewRequest("POST", createTeamURL, bytes.NewBuffer(teamJSON)) + assert.NoError(t, err) + req.Header.Set("Content-Type", "application/json") + req.Header.Set("Authorization", "Bearer "+accessToken) + + resp, err := client.Do(req) + assert.NoError(t, err) + defer resp.Body.Close() + + if resp.StatusCode != http.StatusCreated { + body, _ := io.ReadAll(resp.Body) + t.Fatalf("Failed to create team: status=%d, body=%s", resp.StatusCode, string(body)) + } + + var team map[string]interface{} + err = json.NewDecoder(resp.Body).Decode(&team) + assert.NoError(t, err) + + teamID := getTeamID(team) + + // Step 2: Create invitation via HTTP API + invitationData := map[string]interface{}{ + "user_id": inviteeUserID, + "email": inviteeUserID + "@test.com", // Provide email so GetUser is not needed + "member_type": "user", + "role_id": "user", + "message": "Test invitation", + } + invJSON, _ := json.Marshal(invitationData) + + createInvURL := fmt.Sprintf("%s%s/user/teams/%s/invitations", serverURL, baseURL, teamID) + invReq, err := http.NewRequest("POST", createInvURL, bytes.NewBuffer(invJSON)) + assert.NoError(t, err) + invReq.Header.Set("Content-Type", "application/json") + invReq.Header.Set("Authorization", "Bearer "+accessToken) + + invResp, err := client.Do(invReq) + assert.NoError(t, err) + defer invResp.Body.Close() + + if invResp.StatusCode != http.StatusCreated { + body, _ := io.ReadAll(invResp.Body) + t.Fatalf("Failed to create invitation: status=%d, body=%s", invResp.StatusCode, string(body)) + } + + var invitation map[string]interface{} + err = json.NewDecoder(invResp.Body).Decode(&invitation) + assert.NoError(t, err) + + invitationID, ok := invitation["invitation_id"].(string) + if !ok { + t.Fatalf("invitation_id not found in response") + } + + return teamID, invitationID +} + +// createUserInDB creates a user record directly in the database using userProvider +// Returns the actual user_id created (which may differ from the requested userID) +func createUserInDB(t *testing.T, userID string) string { + userData := map[string]interface{}{ + "user_id": userID, + "name": "Test User " + userID, + "email": userID + "@test.com", + "status": "enabled", + } + + // Create user using userProvider.CreateUser + provider, err := oauth.OAuth.GetUserProvider() + if err != nil { + t.Fatalf("Failed to get user provider: %v", err) + } + + ctx := context.Background() + createdUserID, err := provider.CreateUser(ctx, userData) + if err != nil { + t.Fatalf("Failed to create user via provider: %v", err) + } + + if createdUserID != userID { + t.Logf("Note: Created user ID %s differs from requested %s", createdUserID, userID) + } + + return createdUserID +} + // createTestInvitation creates a test invitation and returns its ID func createTestInvitation(t *testing.T, serverURL, baseURL, accessToken, teamID, userID string) string { return createTestInvitationWithMessage(t, serverURL, baseURL, accessToken, teamID, userID, "Test invitation") diff --git a/openapi/tests/user/team_test.go b/openapi/tests/user/team_test.go index ac992d10..1696ef21 100644 --- a/openapi/tests/user/team_test.go +++ b/openapi/tests/user/team_test.go @@ -116,24 +116,13 @@ func TestTeamList(t *testing.T) { assert.NoError(t, err, "Should read response body") if resp.StatusCode == 200 { - // Parse response as pagination result - var response map[string]interface{} - err = json.Unmarshal(body, &response) - assert.NoError(t, err, "Should parse JSON response") + // Parse response as array (TeamList returns array directly, not paginated) + var teams []interface{} + err = json.Unmarshal(body, &teams) + assert.NoError(t, err, "Should parse JSON response as array") - // Check pagination structure (consistent with other modules) - if data, ok := response["data"]; ok { - assert.IsType(t, []interface{}{}, data, "Should have data array") - } - if total, ok := response["total"]; ok { - assert.IsType(t, float64(0), total, "Should have total count") - } - if page, ok := response["page"]; ok { - assert.IsType(t, float64(0), page, "Should have page number") - } - if pagesize, ok := response["pagesize"]; ok { - assert.IsType(t, float64(0), pagesize, "Should have pagesize") - } + // Verify it's an array + assert.IsType(t, []interface{}{}, teams, "Response should be an array") } t.Logf("Team list test %s: status=%d, body=%s", tc.name, resp.StatusCode, string(body)) @@ -415,16 +404,13 @@ func TestTeamGet(t *testing.T) { if resp.StatusCode == 200 { if tc.teamID == "" { - // Parse response as team list (pagination result) - var response map[string]interface{} - err = json.Unmarshal(body, &response) - assert.NoError(t, err, "Should parse JSON response") + // Parse response as team list (returns array directly, not paginated) + var teams []interface{} + err = json.Unmarshal(body, &teams) + assert.NoError(t, err, "Should parse JSON response as array") - // Check pagination structure - assert.Contains(t, response, "data", "Should have data array") - assert.Contains(t, response, "total", "Should have total count") - assert.Contains(t, response, "page", "Should have page number") - assert.Contains(t, response, "pagesize", "Should have pagesize") + // Verify it's an array + assert.IsType(t, []interface{}{}, teams, "Response should be an array") } else { // Parse response as team detail object var team map[string]interface{} diff --git a/openapi/user/team_invitation.go b/openapi/user/team_invitation.go index 6447ebd0..2f7d21b1 100644 --- a/openapi/user/team_invitation.go +++ b/openapi/user/team_invitation.go @@ -451,6 +451,159 @@ func GinTeamInvitationDelete(c *gin.Context) { response.RespondWithSuccess(c, http.StatusOK, gin.H{"message": "Invitation cancelled successfully"}) } +// GinTeamInvitationAccept handles POST /user/teams/invitations/:invitation_id/accept - Accept invitation and login to team +func GinTeamInvitationAccept(c *gin.Context) { + // Get authorized user info + authInfo := oauth.GetAuthorizedInfo(c) + if authInfo == nil || authInfo.UserID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidClient.Code, + ErrorDescription: "User not authenticated", + } + response.RespondWithError(c, response.StatusUnauthorized, errorResp) + return + } + + ctx := c.Request.Context() + + // Use authInfo.UserID directly - it might be OAuth subject, but LoginByTeamID will handle user creation + userID := authInfo.UserID + + invitationID := c.Param("invitation_id") + if invitationID == "" { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Invitation ID is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Parse request body to get token + var req struct { + Token string `json:"token" binding:"required"` + } + if err := c.ShouldBindJSON(&req); err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Invalid request body: token is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Get user provider instance + provider, err := getUserProvider() + if err != nil { + log.Error("Failed to get user provider: %v", err) + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to process invitation", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Get invitation details first to retrieve team_id + invitationData, err := provider.GetMemberByInvitationID(ctx, invitationID) + if err != nil { + log.Error("Failed to get invitation: %v", err) + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Invitation not found", + } + response.RespondWithError(c, response.StatusNotFound, errorResp) + return + } + + // Get team_id from invitation + teamID := toString(invitationData["team_id"]) + if teamID == "" { + log.Error("Invalid invitation: missing team_id") + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Invalid invitation data", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + 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) + if err != nil { + log.Error("Failed to accept invitation: %v", err) + // Check error type for appropriate response + if strings.Contains(err.Error(), "not found") || strings.Contains(err.Error(), "already accepted") { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Invitation not found or already accepted", + } + response.RespondWithError(c, response.StatusNotFound, errorResp) + } else if strings.Contains(err.Error(), "expired") { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Invitation has expired", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + } else { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to accept invitation", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + } + return + } + + // Prepare login context with full device/platform information + loginCtx := makeLoginContext(c) + + // Login with the team that was just joined + // Note: userID must exist in database (user table) + loginResponse, err := LoginByTeamID(userID, teamID, loginCtx) + if err != nil { + log.Error("Failed to login with team: %v", err) + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Invitation accepted but failed to login: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Revoke the current token if it exists (similar to team selection) + currentToken := oauth.OAuth.GetAccessToken(c) + if currentToken != "" { + if err := oauth.OAuth.Revoke(ctx, currentToken, "access_token"); err != nil { + // Log the error but don't fail the request + log.Warn("Failed to revoke previous token: %v", err) + } + } + + // Send secure cookies (access token, refresh token, and session ID) + SendLoginCookies(c, loginResponse, "") + + // Return the new tokens in response body + response.RespondWithSuccess(c, http.StatusOK, loginResponse) +} + // Yao Process Handlers (for Yao application calls) // ProcessTeamInvitationList user.team.invitation.list Team invitation list processor diff --git a/openapi/user/user.go b/openapi/user/user.go index 7a07a5f5..4394423f 100644 --- a/openapi/user/user.go +++ b/openapi/user/user.go @@ -58,7 +58,8 @@ func Attach(group *gin.RouterGroup, oauth types.OAuth) { func attachTeam(group *gin.RouterGroup, oauth types.OAuth) { // Public endpoint for viewing team invitations (no auth required) // Must be registered BEFORE the team group with auth guard - group.GET("/teams/invitations/:invitation_id", GinTeamInvitationGetPublic) // GET /user/teams/invitations/:invitation_id - Get invitation details (public) + group.GET("/teams/invitations/:invitation_id", GinTeamInvitationGetPublic) // GET /user/teams/invitations/:invitation_id - Get invitation details (public) + group.POST("/teams/invitations/:invitation_id/accept", oauth.Guard, GinTeamInvitationAccept) // POST /user/teams/invitations/:invitation_id/accept - Accept invitation and login team := group.Group("/teams")