Refactor invitation handling and data parsing in user provider

- Updated TestMemberInvitationExpiry to use a consistent expired time for invitations.
- Simplified invitation expiration check in AcceptInvitation by utilizing a new checkTimeExpired function.
- Enhanced member ID parsing in AcceptInvitation to handle potential errors more gracefully.
- Refactored total extraction logic in CountOAuthAccounts, CountRoles, CountTeams, CountTypes, and CountUsers to use a new utility function for improved clarity and consistency.
- Introduced parseTimeFromDB and parseIntFromDB functions to streamline data parsing from database fields.
This commit is contained in:
Max 2025-09-16 17:54:56 +08:00
parent b20ddee701
commit 6c73395158
8 changed files with 106 additions and 127 deletions

View file

@ -228,16 +228,15 @@ func (u *DefaultUser) AcceptInvitation(ctx context.Context, invitationToken stri
member := members[0]
// Check if invitation has expired
if expiresAt, ok := member["invitation_expires_at"]; ok {
if expiryTime, ok := expiresAt.(time.Time); ok {
if time.Now().After(expiryTime) {
return fmt.Errorf("invitation has expired")
}
}
if expired, err := checkTimeExpired(member["invitation_expires_at"]); err == nil && expired {
return fmt.Errorf("invitation has expired")
}
// Update member status to active
memberID := member["id"].(int64)
memberID, err := parseIntFromDB(member["id"])
if err != nil {
return fmt.Errorf("invalid member ID: %w", err)
}
updateData := maps.MapStrAny{
"status": "active",
"joined_at": time.Now(),
@ -489,15 +488,10 @@ func (u *DefaultUser) UpdateMemberLastActivity(ctx context.Context, teamID strin
return err
}
loginCount := 0
loginCount := int64(0)
if count := member["login_count"]; count != nil {
switch v := count.(type) {
case int:
loginCount = v
case int64:
loginCount = int(v)
case int32:
loginCount = int(v)
if parsedCount, err := parseIntFromDB(count); err == nil {
loginCount = parsedCount
}
}
updateData["login_count"] = loginCount + 1

View file

@ -722,6 +722,7 @@ func TestMemberInvitationExpiry(t *testing.T) {
assert.NoError(t, err)
// Create member with expired invitation
expiredTime := time.Now().Add(-2 * time.Hour) // Expired 2 hours ago to be safe
memberData := maps.MapStrAny{
"team_id": teamID,
"user_id": inviteeUser,
@ -729,9 +730,9 @@ func TestMemberInvitationExpiry(t *testing.T) {
"role_id": "user",
"status": "pending",
"invited_by": ownerUser,
"invited_at": time.Now(),
"invited_at": expiredTime.Add(-1 * time.Hour), // Invited 3 hours ago
"invitation_token": "expired-token-" + testUUID,
"invitation_expires_at": time.Now().Add(-1 * time.Hour), // Expired 1 hour ago
"invitation_expires_at": expiredTime, // Expired 2 hours ago
}
_, err = testProvider.CreateMember(ctx, memberData)

View file

@ -215,29 +215,9 @@ func (u *DefaultUser) CountOAuthAccounts(ctx context.Context, param model.QueryP
return 0, fmt.Errorf(ErrFailedToGetOAuthAccount, err)
}
// Extract total from pagination result
if total, ok := result["total"].(int64); ok {
return total, nil
}
// Handle different total types returned by Paginate
// Extract total from pagination result using utility function
if totalInterface, ok := result["total"]; ok {
switch v := totalInterface.(type) {
case int:
return int64(v), nil
case int32:
return int64(v), nil
case int64:
return v, nil
case uint:
return int64(v), nil
case uint32:
return int64(v), nil
case uint64:
return int64(v), nil
default:
return 0, fmt.Errorf("unexpected total type: %T", totalInterface)
}
return parseIntFromDB(totalInterface)
}
return 0, fmt.Errorf("total not found in pagination result")

View file

@ -212,29 +212,9 @@ func (u *DefaultUser) CountRoles(ctx context.Context, param model.QueryParam) (i
return 0, fmt.Errorf(ErrFailedToGetRole, err)
}
// Extract total from pagination result
if total, ok := result["total"].(int64); ok {
return total, nil
}
// Handle different total types returned by Paginate
// Extract total from pagination result using utility function
if totalInterface, ok := result["total"]; ok {
switch v := totalInterface.(type) {
case int:
return int64(v), nil
case int32:
return int64(v), nil
case int64:
return v, nil
case uint:
return int64(v), nil
case uint32:
return int64(v), nil
case uint64:
return int64(v), nil
default:
return 0, fmt.Errorf("unexpected total type: %T", totalInterface)
}
return parseIntFromDB(totalInterface)
}
return 0, fmt.Errorf("total not found in pagination result")

View file

@ -228,29 +228,9 @@ func (u *DefaultUser) CountTeams(ctx context.Context, param model.QueryParam) (i
return 0, fmt.Errorf(ErrFailedToGetTeam, err)
}
// Extract total from pagination result
if total, ok := result["total"].(int64); ok {
return total, nil
}
// Handle different total types returned by Paginate
// Extract total from pagination result using utility function
if totalInterface, ok := result["total"]; ok {
switch v := totalInterface.(type) {
case int:
return int64(v), nil
case int32:
return int64(v), nil
case int64:
return v, nil
case uint:
return int64(v), nil
case uint32:
return int64(v), nil
case uint64:
return int64(v), nil
default:
return 0, fmt.Errorf("unexpected total type: %T", totalInterface)
}
return parseIntFromDB(totalInterface)
}
return 0, fmt.Errorf("total not found in pagination result")

View file

@ -202,29 +202,9 @@ func (u *DefaultUser) CountTypes(ctx context.Context, param model.QueryParam) (i
return 0, fmt.Errorf(ErrFailedToGetType, err)
}
// Extract total from pagination result
if total, ok := result["total"].(int64); ok {
return total, nil
}
// Handle different total types returned by Paginate
// Extract total from pagination result using utility function
if totalInterface, ok := result["total"]; ok {
switch v := totalInterface.(type) {
case int:
return int64(v), nil
case int32:
return int64(v), nil
case int64:
return v, nil
case uint:
return int64(v), nil
case uint32:
return int64(v), nil
case uint64:
return int64(v), nil
default:
return 0, fmt.Errorf("unexpected total type: %T", totalInterface)
}
return parseIntFromDB(totalInterface)
}
return 0, fmt.Errorf("total not found in pagination result")

View file

@ -52,29 +52,9 @@ func (u *DefaultUser) CountUsers(ctx context.Context, param model.QueryParam) (i
return 0, fmt.Errorf(ErrFailedToGetUser, err)
}
// Extract total from pagination result
if total, ok := result["total"].(int64); ok {
return total, nil
}
// Handle different total types returned by Paginate
// Extract total from pagination result using utility function
if totalInterface, ok := result["total"]; ok {
switch v := totalInterface.(type) {
case int:
return int64(v), nil
case int32:
return int64(v), nil
case int64:
return v, nil
case uint:
return int64(v), nil
case uint32:
return int64(v), nil
case uint64:
return int64(v), nil
default:
return 0, fmt.Errorf("unexpected total type: %T", totalInterface)
}
return parseIntFromDB(totalInterface)
}
return 0, fmt.Errorf("total not found in pagination result")

View file

@ -4,6 +4,7 @@ import (
"context"
"crypto/rand"
"fmt"
"time"
"github.com/google/uuid"
gonanoid "github.com/matoous/go-nanoid/v2"
@ -166,3 +167,86 @@ func generateRandomPassword(length int) (string, error) {
return string(bytes), nil
}
// parseTimeFromDB parses time values from database fields, handling different formats and types
func parseTimeFromDB(value interface{}) (*time.Time, error) {
if value == nil {
return nil, nil
}
switch v := value.(type) {
case time.Time:
return &v, nil
case string:
if v == "" {
return nil, nil
}
// Try parsing common time formats - assume local timezone for database timestamps
if parsedTime, err := time.ParseInLocation("2006-01-02 15:04:05", v, time.Local); err == nil {
return &parsedTime, nil
}
if parsedTime, err := time.Parse(time.RFC3339, v); err == nil {
return &parsedTime, nil
}
if parsedTime, err := time.ParseInLocation("2006-01-02T15:04:05", v, time.Local); err == nil {
return &parsedTime, nil
}
if parsedTime, err := time.ParseInLocation("2006-01-02 15:04:05.000000", v, time.Local); err == nil {
return &parsedTime, nil
}
return nil, fmt.Errorf("unable to parse time format: %s", v)
default:
return nil, fmt.Errorf("unsupported time type: %T", value)
}
}
// parseIntFromDB parses integer values from database fields, handling different integer types
func parseIntFromDB(value interface{}) (int64, error) {
if value == nil {
return 0, fmt.Errorf("value is nil")
}
switch v := value.(type) {
case int64:
return v, nil
case int:
return int64(v), nil
case int32:
return int64(v), nil
case uint:
return int64(v), nil
case uint32:
return int64(v), nil
case uint64:
// Check for overflow
if v > 9223372036854775807 { // max int64
return 0, fmt.Errorf("value too large for int64: %d", v)
}
return int64(v), nil
case float64:
// Handle cases where database returns numbers as floats
return int64(v), nil
case string:
// Try to parse string as integer
if parsed, err := fmt.Sscanf(v, "%d", new(int64)); err == nil && parsed == 1 {
var result int64
fmt.Sscanf(v, "%d", &result)
return result, nil
}
return 0, fmt.Errorf("unable to parse string as integer: %s", v)
default:
return 0, fmt.Errorf("unsupported integer type: %T", value)
}
}
// checkTimeExpired checks if a time field from database indicates expiration
func checkTimeExpired(value interface{}) (bool, error) {
parsedTime, err := parseTimeFromDB(value)
if err != nil {
return false, err
}
if parsedTime == nil {
return false, nil // No expiry time set
}
return time.Now().After(*parsedTime), nil
}