diff --git a/openapi/user/entry.go b/openapi/user/entry.go index 4c289c6f..329b25fb 100644 --- a/openapi/user/entry.go +++ b/openapi/user/entry.go @@ -764,17 +764,6 @@ func GinEntryRegister(c *gin.Context) { return } - // Get user provider - userProvider, err := oauth.OAuth.GetUserProvider() - if err != nil { - errorResp := &response.ErrorResponse{ - Code: response.ErrServerError.Code, - ErrorDescription: "Failed to get user provider: " + err.Error(), - } - response.RespondWithError(c, response.StatusInternalServerError, errorResp) - return - } - // Generate name if not provided name := req.Name if name == "" { @@ -821,18 +810,17 @@ func GinEntryRegister(c *gin.Context) { userData["status"] = "active" } - // Create user - userID, err := userProvider.CreateUser(ctx, userData) + // Create user and default team (with rollback on team creation failure) + userID, err := registerUserWithTeam(ctx, userData, req.Locale) if err != nil { - log.Error("Failed to create user: %v", err) + log.Error("Failed to register user: %v", err) errorResp := &response.ErrorResponse{ Code: response.ErrServerError.Code, - ErrorDescription: "Failed to create user: " + err.Error(), + ErrorDescription: err.Error(), } response.RespondWithError(c, response.StatusInternalServerError, errorResp) return } - log.Info("User registered successfully: %s (user_id: %s)", usernameStr, userID) // If auto_login is false and invite not required, return success without tokens diff --git a/openapi/user/login.go b/openapi/user/login.go index c98b99bd..545d8ca1 100644 --- a/openapi/user/login.go +++ b/openapi/user/login.go @@ -11,6 +11,7 @@ import ( "github.com/gin-gonic/gin" "github.com/yaoapp/gou/session" "github.com/yaoapp/kun/log" + "github.com/yaoapp/kun/maps" "github.com/yaoapp/yao/agent/assistant" "github.com/yaoapp/yao/kb" kbapi "github.com/yaoapp/yao/kb/api" @@ -25,6 +26,71 @@ import ( // kbCollectionCreating tracks collections currently being created to avoid duplicate creation var kbCollectionCreating sync.Map +// registerUserWithTeam creates a new user and automatically creates a default team. +// If team creation fails, the user is rolled back (deleted) to ensure data consistency. +// This is the single entry point for all user registration paths (email/mobile, OAuth third-party, etc.). +// +// Parameters: +// - ctx: context for database operations +// - userData: user fields to pass to CreateUser (name, email, status, role_id, type_id, etc.) +// - locale: user's locale for determining default team name (e.g. "zh-cn", "en") +// +// Returns: +// - userID: the created user's ID +// - error: non-nil if user creation or team creation failed (user is rolled back on team failure) +func registerUserWithTeam(ctx context.Context, userData map[string]interface{}, locale string) (string, error) { + userProvider, err := oauth.OAuth.GetUserProvider() + if err != nil { + return "", fmt.Errorf("failed to get user provider: %w", err) + } + + // Create user + userID, err := userProvider.CreateUser(ctx, userData) + if err != nil { + return "", fmt.Errorf("failed to create user: %w", err) + } + + // Auto-create a default team for the new user + // Use "'s Team" / "的团队" format + // Priority: given_name > name (given_name is more natural as display name) + userName := "" + if v, ok := userData["given_name"].(string); ok && v != "" { + userName = v + } else if v, ok := userData["name"].(string); ok && v != "" { + userName = v + } + var defaultTeamName string + if strings.HasPrefix(strings.ToLower(locale), "zh") { + if userName != "" { + defaultTeamName = userName + "的团队" + } else { + defaultTeamName = "我的团队" + } + } else { + if userName != "" { + defaultTeamName = userName + "'s Team" + } else { + defaultTeamName = "My Team" + } + } + teamData := maps.MapStrAny{ + "name": defaultTeamName, + "locale": locale, + } + defaultTeamID, err := teamCreate(ctx, userID, teamData) + if err != nil { + log.Error("Failed to create default team for user %s: %v", userID, err) + // Rollback: delete the created user since a team is required + if delErr := userProvider.DeleteUser(ctx, userID); delErr != nil { + log.Error("Failed to rollback user %s after team creation failure: %v", userID, delErr) + } + return "", fmt.Errorf("registration failed: unable to initialize team: %w", err) + } + + log.Info("User registered: %s, default team: %s", userID, defaultTeamID) + return userID, nil +} + // getCaptcha is the handler for get captcha image for entry (login/register) func getCaptcha(c *gin.Context) { var option captcha.Option = captcha.NewOption() @@ -70,19 +136,17 @@ func LoginThirdParty(providerID string, userinfo *oauthtypes.OIDCUserInfo, login } } - // Check if user exists + // Auto register user if not exists + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + userProvider, err := oauth.OAuth.GetUserProvider() if err != nil { return nil, err } - // Auto register user if not exists - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - var userID string - // Auto register user if not exists if provider.Register != nil && provider.Register.Auto { userID, err = userProvider.GetOAuthUserID(ctx, providerID, userinfo.Sub) if err != nil && err.Error() == user.ErrOAuthAccountNotFound { @@ -103,18 +167,22 @@ func LoginThirdParty(providerID string, userinfo *oauthtypes.OIDCUserInfo, login "status": status, } - // Auto register user - userID, err = userProvider.CreateUser(ctx, userData) + // Register user with default team (with rollback on failure) + userID, err = registerUserWithTeam(ctx, userData, locale) if err != nil { return nil, err } - // Create OAuth account - userData = userinfo.Map() - userData["provider"] = providerID - _, err = userProvider.CreateOAuthAccount(ctx, userID, userData) + // Create OAuth account link + oauthData := userinfo.Map() + oauthData["provider"] = providerID + _, err = userProvider.CreateOAuthAccount(ctx, userID, oauthData) if err != nil { - return nil, err + // Rollback: delete user and team if OAuth account creation fails + if delErr := userProvider.DeleteUser(ctx, userID); delErr != nil { + log.Error("Failed to rollback user %s after OAuth account creation failure: %v", userID, delErr) + } + return nil, fmt.Errorf("failed to create OAuth account: %w", err) } } } @@ -246,8 +314,23 @@ func LoginByUserID(userid string, loginCtx *LoginContext) (*LoginResponse, error return nil, err } - // If user has teams, return team selection status with temporary access token - if numTeams > 0 { + // If user has exactly one team, auto-select it and skip team selection page + if numTeams == 1 { + teams, err := getUserTeams(ctx, userid) + if err == nil && len(teams) == 1 { + teamID := "" + if v, ok := teams[0]["team_id"].(string); ok { + teamID = v + } + if teamID != "" { + return LoginByTeamID(userid, teamID, loginCtx) + } + } + // Fall through to team selection if we couldn't auto-select + } + + // If user has multiple teams, return team selection status with temporary access token + if numTeams > 1 { // Sign temporary access token for Team Selection var teamSelectionExpire int = 10 * 60 // 10 minutes @@ -328,8 +411,9 @@ func LoginByTeamID(userid string, teamID string, loginCtx *LoginContext) (*Login log.Warn("Failed to store user fingerprint: %s", err.Error()) } - // Handle personal account (no team) + // Handle personal account (no team) - deprecated, all users should use teams if teamID == "" || teamID == "personal" { + log.Warn("Personal account login is deprecated. User %s should select a team.", userid) resp, err := issueTokens(ctx, &IssueTokensParams{ UserID: userid, TeamID: "",