mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-09-25 04:19:05 +02:00
feat: migrate one-time and signup tokens to an actor (#1611)
Co-authored-by: Elias Schneider <login@eliasschneider.com>
This commit is contained in:
co-authored by
Elias Schneider
parent
531bb5f0cf
commit
a1b4e1d2b2
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,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"`
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user