mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-08-20 20:06:33 +00:00
502 lines
15 KiB
Go
502 lines
15 KiB
Go
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/apperror"
|
|
"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
|
|
scimSync ScimSyncScheduler
|
|
}
|
|
|
|
func newService(deps Dependencies, actorService *actor.Service) *Service {
|
|
return &Service{
|
|
db: deps.DB,
|
|
actorService: actorService,
|
|
userCreator: deps.UserCreator,
|
|
signer: deps.Signer,
|
|
auditLog: deps.AuditLog,
|
|
scimSync: deps.ScimSync,
|
|
}
|
|
}
|
|
|
|
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{}, "", apperror.OpenSignupDisabled()
|
|
}
|
|
|
|
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{}, "", apperror.TokenInvalidOrExpired()
|
|
}
|
|
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
|
|
}
|
|
if s.scimSync != nil {
|
|
s.scimSync.ScheduleSync(ctx)
|
|
}
|
|
|
|
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()
|
|
}()
|
|
|
|
// We lock the users table to prevent concurrent initial admin setups from racing to create the first user
|
|
// This is only necessary for Postgres, since SQLite serializes all writes anyway
|
|
if err := lockInitialAdminSetup(ctx, tx); err != nil {
|
|
return model.User{}, "", err
|
|
}
|
|
|
|
// Reject setup when a committed user already exists
|
|
setupCompleted, err := s.isInitialAdminSetupCompleted(ctx, tx)
|
|
if err != nil {
|
|
return model.User{}, "", err
|
|
}
|
|
if setupCompleted {
|
|
return model.User{}, "", apperror.SetupAlreadyCompleted()
|
|
}
|
|
|
|
// Build the first user with administrator privileges
|
|
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
|
|
}
|
|
|
|
// Issue the setup session before committing so failures roll back the transaction
|
|
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
|
|
}
|
|
if s.scimSync != nil {
|
|
s.scimSync.ScheduleSync(ctx)
|
|
}
|
|
|
|
return user, token, nil
|
|
}
|
|
|
|
func lockInitialAdminSetup(ctx context.Context, tx *gorm.DB) error {
|
|
if tx.Name() != "postgres" {
|
|
return nil
|
|
}
|
|
|
|
if err := tx.WithContext(ctx).Exec("LOCK TABLE users IN SHARE ROW EXCLUSIVE MODE").Error; err != nil {
|
|
return fmt.Errorf("failed to lock users table for initial admin setup: %w", err)
|
|
}
|
|
|
|
return 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.Sort.Column, params.Sort.Direction)
|
|
|
|
page := max(params.Pagination.Page, 1)
|
|
pageSize := params.Pagination.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)
|
|
})
|
|
}
|