feat: migrate one-time and signup tokens to an actor (#1611)

Co-authored-by: Elias Schneider <login@eliasschneider.com>
This commit is contained in:
Alessandro (Ale) Segala
2026-07-26 15:32:43 +02:00
committed by GitHub
co-authored by Elias Schneider
parent 531bb5f0cf
commit a1b4e1d2b2
37 changed files with 7989 additions and 905 deletions
+217
View File
@@ -0,0 +1,217 @@
package usersignup
import (
"context"
"errors"
"fmt"
"log/slog"
"time"
"github.com/italypaleale/francis/actor"
"github.com/pocket-id/pocket-id/backend/internal/common"
)
// Signup tokens are stored entirely in the actor state store.
// Each token is its own actor, whose actor ID is the token value itself.
// The state is persisted with a TTL equal to the token's lifetime, so it's purged automatically when the token expires (there's no separate cleanup job and no expiration alarm).
// Listing tokens uses ListStates, which only returns states that haven't expired yet.
// SignupTokenActorType is the actor type for the signup token actor
const SignupTokenActorType = "SignupToken"
// Methods exposed by the signup token actor
// Because we cannot invoke an actor while a DB transaction is open (that would deadlock on SQLite), consuming a token is done by invoking the actor first (which atomically validates it and increments its usage count), and only afterwards performing the remaining work.
// On failure, the caller compensates by releasing the token via the "release" method as best-effort.
const (
// SignupTokenMethodCreate stores a new signup token, replacing any existing state
SignupTokenMethodCreate = "create"
// SignupTokenMethodDelete removes a signup token
SignupTokenMethodDelete = "delete"
signupTokenMethodMigrate = "migrate"
signupTokenMethodConsume = "consume"
signupTokenMethodRelease = "release"
)
// signupTokenConsumeStatus is the outcome of a "consume" invocation.
type signupTokenConsumeStatus string
const (
// signupTokenConsumeOK indicates the token was valid and one use has been consumed
signupTokenConsumeOK signupTokenConsumeStatus = "ok"
// signupTokenConsumeNotFound indicates the token doesn't exist (or has expired)
signupTokenConsumeNotFound signupTokenConsumeStatus = "not_found"
// signupTokenConsumeLimitReached indicates the token has no uses left
signupTokenConsumeLimitReached signupTokenConsumeStatus = "limit_reached"
)
// SignupTokenState is the persisted state of a signup token actor.
// The token value itself is the actor's ID, so it isn't repeated here.
type SignupTokenState struct {
ID string
UsageLimit int
UsageCount int
UserGroupIDs []string
CreatedAt time.Time
ExpiresAt time.Time
}
// signupTokenConsumeResponse is the response of the "consume" method
type signupTokenConsumeResponse struct {
Status signupTokenConsumeStatus
// UserGroupIDs is set only when Status is "ok", and contains the groups the new user should join
UserGroupIDs []string
}
// signupTokenActor is the actor that manages a single signup token
type signupTokenActor struct {
log *slog.Logger
client actor.Client[SignupTokenState]
}
// NewSignupTokenActor allocates a new signup token actor
// It satisfies actor.Factory
func NewSignupTokenActor(actorID string, service *actor.Service) actor.Actor {
return &signupTokenActor{
log: slog.With(
slog.String("scope", "actor"),
slog.String("actorType", SignupTokenActorType),
),
client: actor.NewActorClient[SignupTokenState](SignupTokenActorType, actorID, service),
}
}
// Invoke implements actor.ActorInvoke
func (a *signupTokenActor) Invoke(parentCtx context.Context, method string, data actor.Envelope) (any, error) {
switch method {
case SignupTokenMethodCreate:
return nil, a.create(parentCtx, data, false)
case signupTokenMethodMigrate:
return nil, a.create(parentCtx, data, true)
case signupTokenMethodConsume:
return a.consume(parentCtx)
case signupTokenMethodRelease:
return nil, a.release(parentCtx)
case SignupTokenMethodDelete:
return nil, a.delete(parentCtx)
default:
return nil, common.ErrUnsupportedActorMethod{Method: method}
}
}
// create stores the token's state.
// When onlyIfMissing is true the write is skipped if the actor already has state: this is used by the one-time migration of the pre-actor tokens, so a token that has already been migrated is never reset.
func (a *signupTokenActor) create(parentCtx context.Context, data actor.Envelope, onlyIfMissing bool) error {
if data == nil {
return fmt.Errorf("request body is empty for method '%s'", SignupTokenMethodCreate)
}
var state SignupTokenState
err := data.Decode(&state)
if err != nil {
return fmt.Errorf("request body is not valid for method '%s': %w", SignupTokenMethodCreate, err)
}
if onlyIfMissing {
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
defer cancel()
current, err := a.client.GetState(ctx)
if err != nil {
return fmt.Errorf("error retrieving actor state: %w", err)
}
// An empty ID means there's no state yet
if current.ID != "" {
return nil
}
}
return a.setState(parentCtx, state)
}
// consume atomically validates the token and, if it's still usable, records one more use.
func (a *signupTokenActor) consume(parentCtx context.Context) (signupTokenConsumeResponse, error) {
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
defer cancel()
state, err := a.client.GetState(ctx)
if err != nil {
return signupTokenConsumeResponse{}, fmt.Errorf("error retrieving actor state: %w", err)
}
// An empty ID means there's no state: the token doesn't exist (or its state already expired and was purged)
if state.ID == "" || state.ExpiresAt.Before(time.Now()) {
return signupTokenConsumeResponse{
Status: signupTokenConsumeNotFound,
}, nil
}
if state.UsageCount >= state.UsageLimit {
return signupTokenConsumeResponse{
Status: signupTokenConsumeLimitReached,
}, nil
}
// Consume one use of the token
state.UsageCount++
err = a.setState(parentCtx, state)
if err != nil {
return signupTokenConsumeResponse{}, err
}
return signupTokenConsumeResponse{
Status: signupTokenConsumeOK,
UserGroupIDs: state.UserGroupIDs,
}, nil
}
// release reverts the usage count increment performed while consuming the token, to compensate when the signup could not be completed.
func (a *signupTokenActor) release(parentCtx context.Context) error {
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
defer cancel()
state, err := a.client.GetState(ctx)
if err != nil {
return fmt.Errorf("error retrieving actor state: %w", err)
}
// The token is gone (for example, it expired and was purged) or was never consumed: nothing to compensate
if state.ID == "" || state.UsageCount <= 0 {
return nil
}
state.UsageCount--
return a.setState(parentCtx, state)
}
// delete removes the token.
func (a *signupTokenActor) delete(parentCtx context.Context) error {
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
defer cancel()
err := a.client.DeleteState(ctx)
if err != nil && !errors.Is(err, actor.ErrStateNotFound) {
// Deleting a token that doesn't exist (for example, one that expired in the meanwhile) already reaches the desired end state
return fmt.Errorf("error deleting actor state: %w", err)
}
return nil
}
// setState saves the state with a TTL matching the token's remaining lifetime, so it's purged automatically once the token expires.
// Saving is skipped if the token has already expired, since there would be nothing left to store.
func (a *signupTokenActor) setState(parentCtx context.Context, state SignupTokenState) error {
ttl := time.Until(state.ExpiresAt)
if ttl <= 0 {
return nil
}
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
defer cancel()
err := a.client.SetState(ctx, state, &actor.SetStateOpts{
TTL: ttl,
})
if err != nil {
return fmt.Errorf("error saving actor state: %w", err)
}
return nil
}
+173
View File
@@ -0,0 +1,173 @@
package usersignup
import (
"testing"
"time"
"github.com/italypaleale/francis/actor"
"github.com/italypaleale/francis/host/local"
"github.com/stretchr/testify/require"
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
)
// newSignupTokenActorService starts a test actor host with the signup token actor registered and returns its service
func newSignupTokenActorService(t *testing.T) *actor.Service {
t.Helper()
var svc *actor.Service
testutils.NewActorHostForTest(t, func(t *testing.T, h *local.Host) {
err := h.RegisterActor(SignupTokenActorType, NewSignupTokenActor)
require.NoError(t, err)
svc = h.Service()
})
require.NotNil(t, svc)
return svc
}
func createSignupTokenForTest(t *testing.T, svc *actor.Service, token string, state SignupTokenState) {
t.Helper()
_, err := svc.Invoke(t.Context(), SignupTokenActorType, token, SignupTokenMethodCreate, state)
require.NoError(t, err)
}
func consumeSignupTokenForTest(t *testing.T, svc *actor.Service, token string) signupTokenConsumeResponse {
t.Helper()
res, err := svc.Invoke(t.Context(), SignupTokenActorType, token, signupTokenMethodConsume, nil)
require.NoError(t, err)
var out signupTokenConsumeResponse
err = res.Decode(&out)
require.NoError(t, err)
return out
}
// listSignupTokenIDsForTest returns the actor IDs (that is, the token values) of every stored signup token
func listSignupTokenIDsForTest(t *testing.T, svc *actor.Service) []string {
t.Helper()
res, err := svc.ListStates(t.Context(), SignupTokenActorType, nil)
require.NoError(t, err)
ids := make([]string, len(res.States))
for i, st := range res.States {
ids[i] = st.ActorID
}
return ids
}
func TestSignupTokenActorConsume(t *testing.T) {
svc := newSignupTokenActorService(t)
createSignupTokenForTest(t, svc, "token-1", SignupTokenState{
ID: "id-1",
ExpiresAt: time.Now().Add(time.Hour),
UsageLimit: 1,
UserGroupIDs: []string{"group-a", "group-b"},
CreatedAt: time.Now(),
})
// First consume succeeds and returns the token's user groups
res := consumeSignupTokenForTest(t, svc, "token-1")
require.Equal(t, signupTokenConsumeOK, res.Status)
require.Equal(t, []string{"group-a", "group-b"}, res.UserGroupIDs)
// Second consume fails: the usage limit (1) has been reached
res = consumeSignupTokenForTest(t, svc, "token-1")
require.Equal(t, signupTokenConsumeLimitReached, res.Status)
}
func TestSignupTokenActorConsumeNotFound(t *testing.T) {
svc := newSignupTokenActorService(t)
res := consumeSignupTokenForTest(t, svc, "does-not-exist")
require.Equal(t, signupTokenConsumeNotFound, res.Status)
}
// TestSignupTokenActorCreateExpired verifies that a token that has already expired is never stored, since its state TTL would be in the past
func TestSignupTokenActorCreateExpired(t *testing.T) {
svc := newSignupTokenActorService(t)
createSignupTokenForTest(t, svc, "token-expired", SignupTokenState{
ID: "id-expired",
ExpiresAt: time.Now().Add(-time.Minute),
UsageLimit: 1,
CreatedAt: time.Now().Add(-time.Hour),
})
require.Empty(t, listSignupTokenIDsForTest(t, svc))
require.Equal(t, signupTokenConsumeNotFound, consumeSignupTokenForTest(t, svc, "token-expired").Status)
}
func TestSignupTokenActorRelease(t *testing.T) {
svc := newSignupTokenActorService(t)
createSignupTokenForTest(t, svc, "token-2", SignupTokenState{
ID: "id-2",
ExpiresAt: time.Now().Add(time.Hour),
UsageLimit: 2,
CreatedAt: time.Now(),
})
// Consume both uses
require.Equal(t, signupTokenConsumeOK, consumeSignupTokenForTest(t, svc, "token-2").Status)
require.Equal(t, signupTokenConsumeOK, consumeSignupTokenForTest(t, svc, "token-2").Status)
require.Equal(t, signupTokenConsumeLimitReached, consumeSignupTokenForTest(t, svc, "token-2").Status)
// Release one use (compensation)
_, err := svc.Invoke(t.Context(), SignupTokenActorType, "token-2", signupTokenMethodRelease, nil)
require.NoError(t, err)
// Consuming succeeds again now that a use was released
require.Equal(t, signupTokenConsumeOK, consumeSignupTokenForTest(t, svc, "token-2").Status)
}
func TestSignupTokenActorDelete(t *testing.T) {
svc := newSignupTokenActorService(t)
createSignupTokenForTest(t, svc, "token-3", SignupTokenState{
ID: "id-3",
ExpiresAt: time.Now().Add(time.Hour),
UsageLimit: 1,
CreatedAt: time.Now(),
})
require.Equal(t, []string{"token-3"}, listSignupTokenIDsForTest(t, svc))
_, err := svc.Invoke(t.Context(), SignupTokenActorType, "token-3", SignupTokenMethodDelete, nil)
require.NoError(t, err)
require.Empty(t, listSignupTokenIDsForTest(t, svc))
// The token can no longer be consumed
require.Equal(t, signupTokenConsumeNotFound, consumeSignupTokenForTest(t, svc, "token-3").Status)
// Deleting a token that no longer exists is a no-op
_, err = svc.Invoke(t.Context(), SignupTokenActorType, "token-3", SignupTokenMethodDelete, nil)
require.NoError(t, err)
}
// TestSignupTokenActorMigrateDoesNotOverwrite verifies that the one-time migration never resets a token that was already migrated and used since
func TestSignupTokenActorMigrateDoesNotOverwrite(t *testing.T) {
svc := newSignupTokenActorService(t)
state := SignupTokenState{
ID: "id-4",
ExpiresAt: time.Now().Add(time.Hour),
UsageLimit: 2,
CreatedAt: time.Now(),
}
createSignupTokenForTest(t, svc, "token-4", state)
// Use the token once
require.Equal(t, signupTokenConsumeOK, consumeSignupTokenForTest(t, svc, "token-4").Status)
// Re-running the migration must not reset the usage count
_, err := svc.Invoke(t.Context(), SignupTokenActorType, "token-4", signupTokenMethodMigrate, state)
require.NoError(t, err)
// Only one use is left, so a single consume succeeds and the next one doesn't
require.Equal(t, signupTokenConsumeOK, consumeSignupTokenForTest(t, svc, "token-4").Status)
require.Equal(t, signupTokenConsumeLimitReached, consumeSignupTokenForTest(t, svc, "token-4").Status)
}
-19
View File
@@ -1,19 +0,0 @@
package usersignup
import (
"context"
"time"
"gorm.io/gorm"
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
)
// CleanupExpiredSignupTokens deletes signup tokens that have expired
// It returns the number of rows removed
func CleanupExpiredSignupTokens(ctx context.Context, db *gorm.DB) (int64, error) {
st := db.
WithContext(ctx).
Delete(&SignupToken{}, "expires_at < ?", datatype.DateTime(time.Now()))
return st.RowsAffected, st.Error
}
+11 -9
View File
@@ -56,7 +56,8 @@ func (h *handler) signUpInitialAdmin(c *gin.Context) {
}
var input signUpDto
if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil {
err = dto.ShouldBindWithNormalizedJSON(c, &input)
if err != nil {
_ = c.Error(err)
return
}
@@ -68,7 +69,8 @@ func (h *handler) signUpInitialAdmin(c *gin.Context) {
}
var userDto dto.UserDto
if err := dto.MapStruct(user, &userDto); err != nil {
err = dto.MapStruct(user, &userDto)
if err != nil {
_ = c.Error(err)
return
}
@@ -136,7 +138,8 @@ func (h *handler) listSignupTokens(c *gin.Context) {
}
var tokensDto []signupTokenDto
if err := dto.MapStructList(tokens, &tokensDto); err != nil {
err = dto.MapStructList(tokens, &tokensDto)
if err != nil {
_ = c.Error(err)
return
}
@@ -183,15 +186,13 @@ func (h *handler) signup(c *gin.Context) {
}
var input signUpDto
if err := dto.ShouldBindWithNormalizedJSON(c, &input); err != nil {
err = dto.ShouldBindWithNormalizedJSON(c, &input)
if err != nil {
_ = c.Error(err)
return
}
ipAddress := c.ClientIP()
userAgent := c.GetHeader("User-Agent")
user, accessToken, err := h.service.SignUp(c.Request.Context(), config, input, ipAddress, userAgent)
user, accessToken, err := h.service.SignUp(c.Request.Context(), config, input, c.ClientIP(), c.GetHeader("User-Agent"))
if err != nil {
_ = c.Error(err)
return
@@ -201,7 +202,8 @@ func (h *handler) signup(c *gin.Context) {
cookie.AddAccessTokenCookie(c, maxAge, accessToken)
var userDto dto.UserDto
if err := dto.MapStruct(user, &userDto); err != nil {
err = dto.MapStruct(user, &userDto)
if err != nil {
_ = c.Error(err)
return
}
+105
View File
@@ -0,0 +1,105 @@
package usersignup
import (
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
"time"
"gorm.io/gorm"
"github.com/pocket-id/pocket-id/backend/internal/model"
)
// This file holds the one-time migration of the pre-actor signup tokens.
// The "actor tokens" migration freezes the signup_tokens table (and its user-group associations) into a JSON document stored in the "kv" table under the "signup_tokens_migrated" key.
// It's loaded here to create the per-token actors on first startup.
// signupTokensMigratedKey is the kv key under which the pre-actor signup tokens were frozen.
const signupTokensMigratedKey = "signup_tokens_migrated" //nolint:gosec // G101 false positive: this is the name of a kv key, not a credential
// migratedSignupToken is the JSON shape of a signup token frozen into the kv table by the migration.
// All timestamps are expressed as Unix seconds.
type migratedSignupToken struct {
ID string `json:"id"`
Token string `json:"token"`
ExpiresAt int64 `json:"expiresAt"`
UsageLimit int `json:"usageLimit"`
UsageCount int `json:"usageCount"`
UserGroupIDs []string `json:"userGroupIds"`
CreatedAt int64 `json:"createdAt"`
}
// migrateSignupTokens creates an actor for every signup token frozen into the kv table by the migration.
// It requires the actor state store to be available, so it must run after the actor host is ready.
// It is idempotent: tokens that have already been migrated are left untouched, so a token that has been used since it was migrated is never reset.
func (s *Service) migrateSignupTokens(ctx context.Context) error {
migrated, err := loadMigratedSignupTokens(ctx, s.db)
if err != nil {
return err
}
if len(migrated) == 0 {
return nil
}
var count int
for _, m := range migrated {
// Skip tokens that have already expired, since there would be nothing left to store
expiresAt := time.Unix(m.ExpiresAt, 0)
if !expiresAt.After(time.Now()) {
continue
}
state := SignupTokenState{
ID: m.ID,
ExpiresAt: expiresAt,
UsageLimit: m.UsageLimit,
UsageCount: m.UsageCount,
UserGroupIDs: m.UserGroupIDs,
CreatedAt: time.Unix(m.CreatedAt, 0),
}
// The token's value is the actor's ID
// The "migrate" method only writes the state if the actor doesn't have one already
_, err = s.actorService.Invoke(ctx, SignupTokenActorType, m.Token, signupTokenMethodMigrate, state)
if err != nil {
return fmt.Errorf("error migrating signup token '%s': %w", m.ID, err)
}
count++
}
slog.InfoContext(ctx, "Migrated signup tokens to actors", slog.Int("count", count))
return nil
}
// loadMigratedSignupTokens reads the signup tokens frozen into the kv table by the migration
// It returns nil if there's nothing to migrate
func loadMigratedSignupTokens(ctx context.Context, db *gorm.DB) ([]migratedSignupToken, error) {
row := model.KV{
Key: signupTokensMigratedKey,
}
ctx, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
err := db.WithContext(ctx).First(&row).Error
switch {
case errors.Is(err, gorm.ErrRecordNotFound):
// There are no migrated signup tokens in the database, nothing to do
return nil, nil
case err != nil:
return nil, fmt.Errorf("failed to load migrated signup tokens from the database: %w", err)
case row.Value == nil || len(*row.Value) == 0:
// Also no migrated signup tokens, nothing to do
return nil, nil
}
var migrated []migratedSignupToken
err = json.Unmarshal([]byte(*row.Value), &migrated)
if err != nil {
return nil, fmt.Errorf("error parsing migrated signup tokens: %w", err)
}
return migrated, nil
}
@@ -0,0 +1,207 @@
package usersignup
import (
"testing"
"time"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
"github.com/pocket-id/pocket-id/backend/internal/utils"
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
)
// versionBeforeMoveTokens is the migration version right before the "actor tokens" migration.
const versionBeforeMoveTokens = 20260722120000
// seedSignupTokensForMigration seeds two signup tokens (one with a user group, one without) into the pre-migration schema.
func seedSignupTokensForMigration(t *testing.T, db *gorm.DB, createdAt, expiresAt time.Time) {
t.Helper()
// An unrelated, non-JSON kv entry, to ensure the freeze/restore queries don't choke on other kv keys
err := db.Exec(
`INSERT INTO kv ("key", "value") VALUES ('instance_id', ?)`,
"not-json-instance-id",
).Error
require.NoError(t, err)
// A user group referenced by one of the tokens
err = db.Exec(
`INSERT INTO user_groups (id, created_at, friendly_name, name) VALUES (?, ?, ?, ?)`,
"grp-1", createdAt.Unix(), "Group One", "group-one",
).Error
require.NoError(t, err)
// A token with a user group
err = db.Exec(
`INSERT INTO signup_tokens (id, created_at, token, expires_at, usage_limit, usage_count) VALUES (?, ?, ?, ?, ?, ?)`,
"tok-1", createdAt.Unix(), "TOKENWITHGROUP01", expiresAt.Unix(), 3, 1,
).Error
require.NoError(t, err)
err = db.Exec(
`INSERT INTO signup_tokens_user_groups (signup_token_id, user_group_id) VALUES (?, ?)`,
"tok-1", "grp-1",
).Error
require.NoError(t, err)
// A token without user groups
err = db.Exec(
`INSERT INTO signup_tokens (id, created_at, token, expires_at, usage_limit, usage_count) VALUES (?, ?, ?, ?, ?, ?)`,
"tok-2", createdAt.Unix(), "TOKENNOGROUP0002", expiresAt.Unix(), 1, 0,
).Error
require.NoError(t, err)
}
func TestLoadMigratedSignupTokens(t *testing.T) {
createdAt := time.Now().Add(-time.Hour).Truncate(time.Second)
expiresAt := time.Now().Add(24 * time.Hour).Truncate(time.Second)
db := testutils.NewDatabaseForTestWithMigrationSeed(t, versionBeforeMoveTokens, func(t *testing.T, db *gorm.DB) {
seedSignupTokensForMigration(t, db, createdAt, expiresAt)
})
// The migration must have dropped the signup_tokens tables
ok := db.Migrator().HasTable("signup_tokens")
require.False(t, ok, "signup_tokens table should have been dropped")
ok = db.Migrator().HasTable("signup_tokens_user_groups")
require.False(t, ok, "signup_tokens_user_groups table should have been dropped")
tokens, err := loadMigratedSignupTokens(t.Context(), db)
require.NoError(t, err)
require.Len(t, tokens, 2)
byID := make(map[string]migratedSignupToken, len(tokens))
for _, tok := range tokens {
byID[tok.ID] = tok
}
tok1 := byID["tok-1"]
require.Equal(t, "TOKENWITHGROUP01", tok1.Token)
require.Equal(t, 3, tok1.UsageLimit)
require.Equal(t, 1, tok1.UsageCount)
require.Equal(t, []string{"grp-1"}, tok1.UserGroupIDs)
require.Equal(t, expiresAt.Unix(), tok1.ExpiresAt)
require.Equal(t, createdAt.Unix(), tok1.CreatedAt)
tok2 := byID["tok-2"]
require.Equal(t, "TOKENNOGROUP0002", tok2.Token)
require.Equal(t, 1, tok2.UsageLimit)
require.Equal(t, 0, tok2.UsageCount)
require.Empty(t, tok2.UserGroupIDs)
}
// TestMigrateSignupTokens verifies that the frozen signup tokens are turned into per-token actors, and that already-expired ones are skipped
func TestMigrateSignupTokens(t *testing.T) {
createdAt := time.Now().Add(-time.Hour).Truncate(time.Second)
expiresAt := time.Now().Add(24 * time.Hour).Truncate(time.Second)
db := testutils.NewDatabaseForTestWithMigrationSeed(t, versionBeforeMoveTokens, func(t *testing.T, db *gorm.DB) {
seedSignupTokensForMigration(t, db, createdAt, expiresAt)
// A token that has already expired: it must not be migrated
err := db.Exec(
`INSERT INTO signup_tokens (id, created_at, token, expires_at, usage_limit, usage_count) VALUES (?, ?, ?, ?, ?, ?)`,
"tok-expired", createdAt.Unix(), "EXPIREDTOKEN0003", time.Now().Add(-time.Hour).Unix(), 1, 0,
).Error
require.NoError(t, err)
})
svc := newSignupServiceForTest(t, db, fakeUserCreator{})
err := svc.migrateSignupTokens(t.Context())
require.NoError(t, err)
entries, err := svc.listSignupTokenStates(t.Context())
require.NoError(t, err)
require.Len(t, entries, 2)
byToken := make(map[string]SignupTokenState, len(entries))
for _, e := range entries {
byToken[e.Token] = e.State
}
tok1 := byToken["TOKENWITHGROUP01"]
require.Equal(t, "tok-1", tok1.ID)
require.Equal(t, 3, tok1.UsageLimit)
require.Equal(t, 1, tok1.UsageCount)
require.Equal(t, []string{"grp-1"}, tok1.UserGroupIDs)
require.Equal(t, expiresAt.Unix(), tok1.ExpiresAt.Unix())
require.Equal(t, createdAt.Unix(), tok1.CreatedAt.Unix())
tok2 := byToken["TOKENNOGROUP0002"]
require.Equal(t, "tok-2", tok2.ID)
require.Empty(t, tok2.UserGroupIDs)
// The expired token must not have been migrated
require.NotContains(t, byToken, "EXPIREDTOKEN0003")
// The migration is idempotent: re-running it doesn't reset a token that has been used since
require.Equal(t, signupTokenConsumeOK, consumeSignupTokenForTest(t, svc.actorService, "TOKENNOGROUP0002").Status)
err = svc.migrateSignupTokens(t.Context())
require.NoError(t, err)
require.Equal(t, signupTokenConsumeLimitReached, consumeSignupTokenForTest(t, svc.actorService, "TOKENNOGROUP0002").Status)
}
// TestLoadMigratedSignupTokensEmpty verifies that when there were no signup tokens, nothing is frozen and nothing is loaded.
func TestLoadMigratedSignupTokensEmpty(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
tokens, err := loadMigratedSignupTokens(t.Context(), db)
require.NoError(t, err)
require.Empty(t, tokens)
}
// TestMoveTokensToActorStateDown verifies that rolling the migration back recreates the signup token tables and restores their contents from the frozen kv document.
func TestMoveTokensToActorStateDown(t *testing.T) {
createdAt := time.Now().Add(-time.Hour).Truncate(time.Second)
expiresAt := time.Now().Add(24 * time.Hour).Truncate(time.Second)
db := testutils.NewDatabaseForTestWithMigrationSeed(t, versionBeforeMoveTokens, func(t *testing.T, db *gorm.DB) {
seedSignupTokensForMigration(t, db, createdAt, expiresAt)
})
// The tables were frozen and dropped by the up migration
ok := db.Migrator().HasTable("signup_tokens")
require.False(t, ok)
// Roll the migration back
sqlDB, err := db.DB()
require.NoError(t, err)
m, cleanup, err := utils.GetEmbeddedMigrateInstance(t.Context(), sqlDB)
require.NoError(t, err)
defer cleanup()
err = m.Migrate(versionBeforeMoveTokens)
require.NoError(t, err)
// The tables must have been recreated and repopulated from the frozen document
ok = db.Migrator().HasTable("signup_tokens")
require.True(t, ok)
ok = db.Migrator().HasTable("signup_tokens_user_groups")
require.True(t, ok)
type row struct {
ID string
Token string
UsageLimit int
UsageCount int
}
var rows []row
err = db.Raw(`SELECT id, token, usage_limit, usage_count FROM signup_tokens ORDER BY id`).Scan(&rows).Error
require.NoError(t, err)
require.Equal(t, []row{
{ID: "tok-1", Token: "TOKENWITHGROUP01", UsageLimit: 3, UsageCount: 1},
{ID: "tok-2", Token: "TOKENNOGROUP0002", UsageLimit: 1, UsageCount: 0},
}, rows)
var groupID string
err = db.Raw(`SELECT user_group_id FROM signup_tokens_user_groups WHERE signup_token_id = ?`, "tok-1").Scan(&groupID).Error
require.NoError(t, err)
require.Equal(t, "grp-1", groupID)
// The frozen document must have been removed from the kv table
var kvCount int64
err = db.Raw(`SELECT count(*) FROM kv WHERE "key" = ?`, signupTokensMigratedKey).Scan(&kvCount).Error
require.NoError(t, err)
require.Zero(t, kvCount)
}
+1 -15
View File
@@ -1,8 +1,6 @@
package usersignup
import (
"time"
"github.com/pocket-id/pocket-id/backend/internal/model"
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
)
@@ -15,17 +13,5 @@ type SignupToken struct {
ExpiresAt datatype.DateTime `json:"expiresAt" sortable:"true"`
UsageLimit int `json:"usageLimit" sortable:"true"`
UsageCount int `json:"usageCount" sortable:"true"`
UserGroups []model.UserGroup `gorm:"many2many:signup_tokens_user_groups;"`
}
func (st *SignupToken) IsExpired() bool {
return time.Time(st.ExpiresAt).Before(time.Now())
}
func (st *SignupToken) IsUsageLimitReached() bool {
return st.UsageCount >= st.UsageLimit
}
func (st *SignupToken) IsValid() bool {
return !st.IsExpired() && !st.IsUsageLimitReached()
UserGroups []model.UserGroup `json:"userGroups"`
}
+26 -3
View File
@@ -2,9 +2,11 @@ package usersignup
import (
"context"
"fmt"
"time"
"github.com/gin-gonic/gin"
"github.com/italypaleale/francis/host/local"
"gorm.io/gorm"
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
@@ -30,7 +32,8 @@ type AppConfigResolver interface {
}
type Dependencies struct {
DB *gorm.DB
DB *gorm.DB
Actors *local.Host
Signer TokenService
AuditLog AuditLogger
@@ -43,12 +46,32 @@ type Module struct {
handler *handler
}
func New(deps Dependencies) *Module {
service := newService(deps)
func New(deps Dependencies) (*Module, error) {
// Register the actor that manages a signup token
// Each token is its own actor, whose actor ID is the token's value
err := deps.Actors.RegisterActor(SignupTokenActorType, NewSignupTokenActor)
if err != nil {
return nil, fmt.Errorf("error registering the %s actor: %w", SignupTokenActorType, err)
}
service := newService(deps, deps.Actors.Service())
return &Module{
service: service,
handler: newHandler(service, deps.AppConfig),
}, nil
}
// RunSignupTokenMigration performs the one-time migration of the pre-actor signup tokens, then blocks until the context is canceled.
// It's meant to be started as a background service gated on the actor host being ready, since the migration needs the actor state store.
// Note that it must not return before the context is canceled, as the service runner stops the application as soon as any of its services returns.
func (m *Module) RunSignupTokenMigration(ctx context.Context) error {
err := m.service.migrateSignupTokens(ctx)
if err != nil {
return fmt.Errorf("failed to migrate signup tokens: %w", err)
}
<-ctx.Done()
return ctx.Err()
}
// RegisterRoutes mounts the signup and signup-token management endpoints
+313 -73
View File
@@ -2,12 +2,15 @@ package usersignup
import (
"context"
"errors"
"fmt"
"log/slog"
"sort"
"strings"
"time"
"github.com/google/uuid"
"github.com/italypaleale/francis/actor"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
"github.com/pocket-id/pocket-id/backend/internal/common"
@@ -22,57 +25,49 @@ import (
const authenticationMethodOneTimePassword = "otp"
type Service struct {
db *gorm.DB
userCreator UserCreator
signer TokenService
auditLog AuditLogger
db *gorm.DB
actorService *actor.Service
userCreator UserCreator
signer TokenService
auditLog AuditLogger
}
func newService(deps Dependencies) *Service {
func newService(deps Dependencies, actorService *actor.Service) *Service {
return &Service{
db: deps.DB,
userCreator: deps.UserCreator,
signer: deps.Signer,
auditLog: deps.AuditLog,
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) {
tx := s.db.Begin()
defer func() {
tx.Rollback()
}()
tokenProvided := signupData.Token != ""
if config.AllowUserSignups.String() != "open" && !tokenProvided {
return model.User{}, "", &common.OpenSignupDisabledError{}
}
var signupToken SignupToken
var userGroupIDs []string
if tokenProvided {
err := tx.
WithContext(ctx).
Preload("UserGroups").
Where("token = ?", signupData.Token).
Clauses(clause.Locking{Strength: "UPDATE"}).
First(&signupToken).
Error
// 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 {
if errors.Is(err, gorm.ErrRecordNotFound) {
return model.User{}, "", &common.TokenInvalidOrExpiredError{}
}
return model.User{}, "", err
return model.User{}, "", fmt.Errorf("error invoking signup token actor: %w", err)
}
if !signupToken.IsValid() {
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{}
}
for _, group := range signupToken.UserGroups {
userGroupIDs = append(userGroupIDs, group.ID)
}
userGroupIDs = consumeRes.UserGroupIDs
}
userToCreate := dto.UserCreateDto{
@@ -85,6 +80,27 @@ func (s *Service) SignUp(ctx context.Context, config *appconfig.AppConfigModel,
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
@@ -97,15 +113,8 @@ func (s *Service) SignUp(ctx context.Context, config *appconfig.AppConfigModel,
if tokenProvided {
s.auditLog.Create(ctx, model.AuditLogEventAccountCreated, ipAddress, userAgent, user.ID, model.AuditLogData{
"signupToken": signupToken.Token,
"signupToken": token,
}, tx)
signupToken.UsageCount++
err = tx.WithContext(ctx).Save(&signupToken).Error
if err != nil {
return model.User{}, "", err
}
} else {
s.auditLog.Create(ctx, model.AuditLogEventAccountCreated, ipAddress, userAgent, user.ID, model.AuditLogData{
"method": "open_signup",
@@ -120,6 +129,19 @@ func (s *Service) SignUp(ctx context.Context, config *appconfig.AppConfigModel,
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() {
@@ -177,55 +199,273 @@ func (s *Service) isInitialAdminSetupCompleted(ctx context.Context, db *gorm.DB)
}
func (s *Service) ListSignupTokens(ctx context.Context, listRequestOptions utils.ListRequestOptions) ([]SignupToken, utils.PaginationResponse, error) {
var tokens []SignupToken
query := s.db.WithContext(ctx).Preload("UserGroups").Model(&SignupToken{})
// 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
}
pagination, err := utils.PaginateFilterAndSort(listRequestOptions, query, &tokens)
return tokens, pagination, 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 {
return s.db.WithContext(ctx).Delete(&SignupToken{}, "id = ?", tokenID).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) {
signupToken, err := newSignupToken(ttl, usageLimit)
if err != nil {
return SignupToken{}, err
}
// Load the referenced user groups to validate them and to include them in the response
var userGroups []model.UserGroup
err = s.db.WithContext(ctx).
Where("id IN ?", userGroupIDs).
Find(&userGroups).
Error
if err != nil {
return SignupToken{}, err
}
signupToken.UserGroups = userGroups
err = s.db.WithContext(ctx).Create(signupToken).Error
if err != nil {
return SignupToken{}, err
if len(userGroupIDs) > 0 {
err := s.db.WithContext(ctx).
Where("id IN ?", userGroupIDs).
Find(&userGroups).
Error
if err != nil {
return SignupToken{}, err
}
}
return *signupToken, nil
}
validGroupIDs := make([]string, len(userGroups))
for i, g := range userGroups {
validGroupIDs[i] = g.ID
}
func newSignupToken(ttl time.Duration, usageLimit int) (*SignupToken, error) {
// 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
}
now := time.Now().Round(time.Second)
token := &SignupToken{
Token: randomString,
ExpiresAt: datatype.DateTime(now.Add(ttl)),
UsageLimit: usageLimit,
UsageCount: 0,
byID := make(map[string]model.UserGroup, len(groups))
for _, g := range groups {
byID[g.ID] = g
}
return token, nil
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)
})
}
+127
View File
@@ -0,0 +1,127 @@
package usersignup
import (
"context"
"errors"
"testing"
"time"
"github.com/stretchr/testify/require"
"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"
"github.com/pocket-id/pocket-id/backend/internal/utils"
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
)
type fakeUserCreator struct {
err error
user model.User
}
func (f fakeUserCreator) CreateUserInternal(_ context.Context, _ *appconfig.AppConfigModel, _ dto.UserCreateDto, _ bool, _ *gorm.DB) (model.User, error) {
if f.err != nil {
return model.User{}, f.err
}
return f.user, nil
}
type fakeSigner struct{}
func (fakeSigner) GenerateAccessToken(_ model.User, _ string, _ time.Duration) (string, error) {
return "access-token", nil
}
type fakeAuditLogger struct{}
func (fakeAuditLogger) Create(_ context.Context, _ model.AuditLogEvent, _, _, _ string, _ model.AuditLogData, _ *gorm.DB) (model.AuditLog, bool) {
return model.AuditLog{}, true
}
func newSignupServiceForTest(t *testing.T, db *gorm.DB, userCreator UserCreator) *Service {
t.Helper()
actorService := newSignupTokenActorService(t)
return newService(Dependencies{
DB: db,
UserCreator: userCreator,
Signer: fakeSigner{},
AuditLog: fakeAuditLogger{},
}, actorService)
}
func signupTokenUsageCount(t *testing.T, svc *Service, tokenID string) int {
t.Helper()
tokens, _, err := svc.ListSignupTokens(t.Context(), listAllOptions())
require.NoError(t, err)
for _, tok := range tokens {
if tok.ID == tokenID {
return tok.UsageCount
}
}
t.Fatalf("signup token %q not found", tokenID)
return 0
}
func TestSignUpConsumesTokenOnSuccess(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
svc := newSignupServiceForTest(t, db, fakeUserCreator{user: model.User{Base: model.Base{ID: "new-user"}}})
token, err := svc.CreateSignupToken(t.Context(), time.Hour, 2, nil)
require.NoError(t, err)
config := appconfig.NewTestConfig(nil)
user, accessToken, err := svc.SignUp(t.Context(), config, signUpDto{
Username: "newuser",
Token: token.Token,
}, "1.2.3.4", "test-agent")
require.NoError(t, err)
require.Equal(t, "new-user", user.ID)
require.Equal(t, "access-token", accessToken)
// The token's usage count must have been incremented and not rolled back
require.Equal(t, 1, signupTokenUsageCount(t, svc, token.ID))
}
func TestSignUpCompensatesTokenOnFailure(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
boom := errors.New("could not create user")
svc := newSignupServiceForTest(t, db, fakeUserCreator{err: boom})
token, err := svc.CreateSignupToken(t.Context(), time.Hour, 2, nil)
require.NoError(t, err)
config := appconfig.NewTestConfig(nil)
_, _, err = svc.SignUp(t.Context(), config, signUpDto{
Username: "newuser",
Token: token.Token,
}, "1.2.3.4", "test-agent")
require.ErrorIs(t, err, boom)
// The usage count increment must have been compensated (reverted back to 0)
require.Equal(t, 0, signupTokenUsageCount(t, svc, token.ID))
}
func TestSignUpRejectsInvalidToken(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
svc := newSignupServiceForTest(t, db, fakeUserCreator{user: model.User{Base: model.Base{ID: "new-user"}}})
config := appconfig.NewTestConfig(nil)
_, _, err := svc.SignUp(t.Context(), config, signUpDto{
Username: "newuser",
Token: "not-a-real-token",
}, "1.2.3.4", "test-agent")
var invalidErr *common.TokenInvalidOrExpiredError
require.ErrorAs(t, err, &invalidErr)
}
// listAllOptions returns list options that return every token on a single page.
func listAllOptions() utils.ListRequestOptions {
var opts utils.ListRequestOptions
opts.Pagination.Page = 1
opts.Pagination.Limit = 100
return opts
}