package usersignup import ( "context" "fmt" "log/slog" "sort" "strings" "time" "github.com/google/uuid" "github.com/italypaleale/francis/actor" "gorm.io/gorm" "github.com/pocket-id/pocket-id/backend/internal/appconfig" "github.com/pocket-id/pocket-id/backend/internal/common" "github.com/pocket-id/pocket-id/backend/internal/dto" "github.com/pocket-id/pocket-id/backend/internal/model" datatype "github.com/pocket-id/pocket-id/backend/internal/model/types" "github.com/pocket-id/pocket-id/backend/internal/utils" ) // authenticationMethodOneTimePassword identifies one-time-password authentication, used for the initial admin setup token // It must match the value emitted by the JWT service in the access token's "amr" claim const authenticationMethodOneTimePassword = "otp" type Service struct { db *gorm.DB actorService *actor.Service userCreator UserCreator signer TokenService auditLog AuditLogger } func newService(deps Dependencies, actorService *actor.Service) *Service { return &Service{ db: deps.DB, actorService: actorService, userCreator: deps.UserCreator, signer: deps.Signer, auditLog: deps.AuditLog, } } func (s *Service) SignUp(ctx context.Context, config *appconfig.AppConfigModel, signupData signUpDto, ipAddress, userAgent string) (model.User, string, error) { tokenProvided := signupData.Token != "" if config.AllowUserSignups.String() != "open" && !tokenProvided { return model.User{}, "", &common.OpenSignupDisabledError{} } var userGroupIDs []string if tokenProvided { // Consume the signup token by invoking its actor: this atomically validates it and increments its usage count // Note: must invoke outside of a DB transaction, since invoking an actor while a transaction is open would deadlock on SQLite res, err := s.actorService.Invoke(ctx, SignupTokenActorType, signupData.Token, signupTokenMethodConsume, nil) if err != nil { return model.User{}, "", fmt.Errorf("error invoking signup token actor: %w", err) } var consumeRes signupTokenConsumeResponse err = res.Decode(&consumeRes) if err != nil { return model.User{}, "", fmt.Errorf("error decoding signup token actor response: %w", err) } if consumeRes.Status != signupTokenConsumeOK { return model.User{}, "", &common.TokenInvalidOrExpiredError{} } userGroupIDs = consumeRes.UserGroupIDs } userToCreate := dto.UserCreateDto{ Username: signupData.Username, Email: signupData.Email, FirstName: signupData.FirstName, LastName: signupData.LastName, DisplayName: strings.TrimSpace(signupData.FirstName + " " + signupData.LastName), UserGroupIds: userGroupIDs, EmailVerified: config.EmailsVerified.IsTrue(), } // The token has now been consumed // From this point on, if we hit an error we compensate by releasing the token (best-effort) user, accessToken, err := s.createSignedUpUser(ctx, config, userToCreate, signupData.Token, tokenProvided, ipAddress, userAgent) if err != nil { if tokenProvided { s.releaseSignupToken(ctx, signupData.Token) } return model.User{}, "", err } return user, accessToken, nil } // createSignedUpUser creates the user and issues an access token within a single transaction. // It performs no actor calls, so it's safe to keep the transaction open for its whole duration. func (s *Service) createSignedUpUser(ctx context.Context, config *appconfig.AppConfigModel, userToCreate dto.UserCreateDto, token string, tokenProvided bool, ipAddress, userAgent string) (model.User, string, error) { tx := s.db.Begin() defer func() { tx.Rollback() }() user, err := s.userCreator.CreateUserInternal(ctx, config, userToCreate, false, tx) if err != nil { return model.User{}, "", err } accessToken, err := s.signer.GenerateAccessToken(user, "", config.SessionDuration.AsDurationMinutes()) if err != nil { return model.User{}, "", err } if tokenProvided { s.auditLog.Create(ctx, model.AuditLogEventAccountCreated, ipAddress, userAgent, user.ID, model.AuditLogData{ "signupToken": token, }, tx) } else { s.auditLog.Create(ctx, model.AuditLogEventAccountCreated, ipAddress, userAgent, user.ID, model.AuditLogData{ "method": "open_signup", }, tx) } err = tx.Commit().Error if err != nil { return model.User{}, "", err } return user, accessToken, nil } // releaseSignupToken reverts the usage count increment performed while consuming a token, used to compensate when the signup could not be completed. // It's a best-effort compensation: if it fails (or the process crashes before it runs) we accept that a token use was consumed unnecessarily func (s *Service) releaseSignupToken(parentCtx context.Context, token string) { // Use a context that is not canceled when the original request ends ctx, cancel := context.WithTimeout(context.WithoutCancel(parentCtx), 10*time.Second) defer cancel() _, err := s.actorService.Invoke(ctx, SignupTokenActorType, token, signupTokenMethodRelease, nil) if err != nil { slog.ErrorContext(ctx, "Failed to release signup token after a failed signup", slog.Any("error", err)) } } func (s *Service) SignUpInitialAdmin(ctx context.Context, config *appconfig.AppConfigModel, signUpData signUpDto) (model.User, string, error) { tx := s.db.Begin() defer func() { tx.Rollback() }() setupCompleted, err := s.isInitialAdminSetupCompleted(ctx, tx) if err != nil { return model.User{}, "", err } if setupCompleted { return model.User{}, "", &common.SetupNotAvailableError{} } userToCreate := dto.UserCreateDto{ FirstName: signUpData.FirstName, LastName: signUpData.LastName, DisplayName: strings.TrimSpace(signUpData.FirstName + " " + signUpData.LastName), Username: signUpData.Username, Email: signUpData.Email, IsAdmin: true, } user, err := s.userCreator.CreateUserInternal(ctx, config, userToCreate, false, tx) if err != nil { return model.User{}, "", err } token, err := s.signer.GenerateAccessToken(user, authenticationMethodOneTimePassword, config.SessionDuration.AsDurationMinutes()) if err != nil { return model.User{}, "", err } err = tx.Commit().Error if err != nil { return model.User{}, "", err } return user, token, nil } func (s *Service) IsInitialAdminSetupCompleted(ctx context.Context) (bool, error) { return s.isInitialAdminSetupCompleted(ctx, s.db) } func (s *Service) isInitialAdminSetupCompleted(ctx context.Context, db *gorm.DB) (bool, error) { var userCount int64 if err := db.WithContext(ctx).Model(&model.User{}). Where("id != ?", common.StaticApiKeyUserID). Count(&userCount).Error; err != nil { return false, err } return userCount != 0, nil } func (s *Service) ListSignupTokens(ctx context.Context, listRequestOptions utils.ListRequestOptions) ([]SignupToken, utils.PaginationResponse, error) { // Each signup token is its own actor, so we enumerate the stored states (expired ones are filtered out by the state store), then sort and paginate in memory entries, err := s.listSignupTokenStates(ctx) if err != nil { return nil, utils.PaginationResponse{}, err } // Resolve the referenced user groups so they can be included in the response groupsByID, err := s.loadUserGroupsByID(ctx, entries) if err != nil { return nil, utils.PaginationResponse{}, err } tokens := make([]SignupToken, len(entries)) for i, e := range entries { tokens[i] = signupTokenModelFromState(e.Token, e.State, resolveUserGroups(e.State.UserGroupIDs, groupsByID)) } return paginateSignupTokens(tokens, listRequestOptions) } func (s *Service) DeleteSignupToken(ctx context.Context, tokenID string) error { // Tokens are addressed by their value (the actor ID), while the API deletes them by ID, so we look up the matching token first entries, err := s.listSignupTokenStates(ctx) if err != nil { return err } for _, e := range entries { if e.State.ID != tokenID { continue } _, err = s.actorService.Invoke(ctx, SignupTokenActorType, e.Token, SignupTokenMethodDelete, nil) if err != nil { return fmt.Errorf("error deleting signup token via actor: %w", err) } return nil } // The token doesn't exist (or has expired): deleting it already reaches the desired end state return nil } // signupTokenEntry pairs a signup token's value (which is its actor ID) with its stored state type signupTokenEntry struct { Token string State SignupTokenState } // listSignupTokenStates returns every signup token currently stored in the actor state store. // Expired tokens are not returned, since the state store filters out states whose TTL has passed. func (s *Service) listSignupTokenStates(ctx context.Context) ([]signupTokenEntry, error) { var ( entries []signupTokenEntry after string ) for { res, err := s.actorService.ListStates(ctx, SignupTokenActorType, &actor.ListStatesOpts{ IncludeData: true, After: after, }) if err != nil { return nil, fmt.Errorf("error listing signup token states: %w", err) } for _, st := range res.States { if st.Data == nil { continue } var state SignupTokenState err = st.Data.Decode(&state) if err != nil { return nil, fmt.Errorf("error decoding state of signup token actor '%s': %w", st.ActorID, err) } entries = append(entries, signupTokenEntry{ Token: st.ActorID, State: state, }) } // An empty cursor means we've just read the last page after = res.AfterID() if after == "" { break } } return entries, nil } func (s *Service) CreateSignupToken(ctx context.Context, ttl time.Duration, usageLimit int, userGroupIDs []string) (SignupToken, error) { // Load the referenced user groups to validate them and to include them in the response var userGroups []model.UserGroup if len(userGroupIDs) > 0 { err := s.db.WithContext(ctx). Where("id IN ?", userGroupIDs). Find(&userGroups). Error if err != nil { return SignupToken{}, err } } validGroupIDs := make([]string, len(userGroups)) for i, g := range userGroups { validGroupIDs[i] = g.ID } // Generate a random token randomString, err := utils.GenerateRandomAlphanumericString(16) if err != nil { return SignupToken{}, err } now := time.Now().Round(time.Second) state := SignupTokenState{ ID: uuid.NewString(), ExpiresAt: now.Add(ttl), UsageLimit: usageLimit, UsageCount: 0, UserGroupIDs: validGroupIDs, CreatedAt: now, } // The token's value is the actor's ID _, err = s.actorService.Invoke(ctx, SignupTokenActorType, randomString, SignupTokenMethodCreate, state) if err != nil { return SignupToken{}, fmt.Errorf("error creating signup token via actor: %w", err) } return signupTokenModelFromState(randomString, state, userGroups), nil } // loadUserGroupsByID loads every user group referenced by the given tokens, keyed by ID. func (s *Service) loadUserGroupsByID(ctx context.Context, entries []signupTokenEntry) (map[string]model.UserGroup, error) { idSet := make(map[string]struct{}) for _, e := range entries { for _, id := range e.State.UserGroupIDs { idSet[id] = struct{}{} } } if len(idSet) == 0 { return map[string]model.UserGroup{}, nil } ids := make([]string, 0, len(idSet)) for id := range idSet { ids = append(ids, id) } var groups []model.UserGroup err := s.db.WithContext(ctx). Where("id IN ?", ids). Find(&groups). Error if err != nil { return nil, err } byID := make(map[string]model.UserGroup, len(groups)) for _, g := range groups { byID[g.ID] = g } return byID, nil } // resolveUserGroups maps the given group IDs to the corresponding UserGroup objects, preserving order and skipping any that no longer exist. func resolveUserGroups(ids []string, byID map[string]model.UserGroup) []model.UserGroup { if len(ids) == 0 { return nil } groups := make([]model.UserGroup, 0, len(ids)) for _, id := range ids { g, ok := byID[id] if ok { groups = append(groups, g) } } return groups } // signupTokenModelFromState builds the API/model representation of a signup token from its actor ID (the token's value) and stored state. func signupTokenModelFromState(token string, state SignupTokenState, groups []model.UserGroup) SignupToken { return SignupToken{ Base: model.Base{ ID: state.ID, CreatedAt: datatype.DateTime(state.CreatedAt), }, Token: token, ExpiresAt: datatype.DateTime(state.ExpiresAt), UsageLimit: state.UsageLimit, UsageCount: state.UsageCount, UserGroups: groups, } } // paginateSignupTokens sorts and paginates the in-memory list of signup tokens, mirroring the behavior of the DB-backed pagination utility. func paginateSignupTokens(tokens []SignupToken, params utils.ListRequestOptions) ([]SignupToken, utils.PaginationResponse, error) { sortSignupTokens(tokens, params.SortColumn, params.SortDirection) page := max(params.Page, 1) pageSize := params.Limit switch { case pageSize < 1: pageSize = 20 case pageSize > 100: pageSize = 100 } totalItems := int64(len(tokens)) totalPages := (totalItems + int64(pageSize) - 1) / int64(pageSize) if totalItems == 0 { totalPages = 1 } if int64(page) > totalPages { page = int(totalPages) } start := min((page-1)*pageSize, len(tokens)) end := min(start+pageSize, len(tokens)) return tokens[start:end], utils.PaginationResponse{ TotalPages: totalPages, TotalItems: totalItems, CurrentPage: page, ItemsPerPage: pageSize, }, nil } // sortSignupTokens sorts the tokens by the given column and direction. // It defaults to sorting by creation date ascending, matching the DB-backed listing. func sortSignupTokens(tokens []SignupToken, column, direction string) { desc := utils.NormalizeSortDirection(direction) == "desc" less := func(i, j int) bool { caI := time.Time(tokens[i].CreatedAt) caJ := time.Time(tokens[j].CreatedAt) return caI.Before(caJ) } switch column { case "expiresAt": less = func(i, j int) bool { eaI := time.Time(tokens[i].ExpiresAt) eaJ := time.Time(tokens[j].ExpiresAt) return eaI.Before(eaJ) } case "usageLimit": less = func(i, j int) bool { return tokens[i].UsageLimit < tokens[j].UsageLimit } case "usageCount": less = func(i, j int) bool { return tokens[i].UsageCount < tokens[j].UsageCount } case "createdAt", "": // Use the default comparator (creation date) default: // Unknown or non-sortable column: keep the default (creation date) ordering } sort.SliceStable(tokens, func(i, j int) bool { if desc { return less(j, i) } return less(i, j) }) }