Files

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)
})
}