mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-07-22 21:35:22 +03:00
Compare commits
5 Commits
main
...
actors/tok
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c572b55598 | ||
|
|
282c68a97a | ||
|
|
09c4b116e6 | ||
|
|
3a18393a4e | ||
|
|
a34203fa88 |
@@ -98,6 +98,34 @@ func (o *NewActorsOpts) getPSK() ([]byte, error) {
|
||||
return crypto.DeriveKey(o.EnvConfig.EncryptionKey, "pocketid/actors-psk/"+o.InstanceID)
|
||||
}
|
||||
|
||||
// NewActorStateStore creates a minimal actor host that can read and write actor state directly, without joining the cluster or binding a network port.
|
||||
// It's meant for short-lived contexts such as CLI commands that need to persist actor state (for example, one-time access tokens) without running the full actor host.
|
||||
// The returned host must NOT be Run(): only direct state operations (Get/Set/Delete on state) are supported, and they require the actor state tables to already exist, which is the case whenever the server has run at least once against this database.
|
||||
func NewActorStateStore(db *gorm.DB, pg *pgxpool.Pool) (*local.Host, error) {
|
||||
opts := &NewActorsOpts{DB: db, Postgres: pg}
|
||||
if pg == nil {
|
||||
sqlDB, err := db.DB()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get *sql.DB connection from Gorm: %w", err)
|
||||
}
|
||||
opts.SQLite = sqlDB
|
||||
}
|
||||
|
||||
providerOpt, err := opts.getProvider()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return local.NewHost(
|
||||
// The address is required by the host but never bound, since the host is not Run
|
||||
local.WithAddress("127.0.0.1:1"),
|
||||
local.WithLogger(slog.Default().With("scope", "actor-state-store")),
|
||||
// The health-check deadline only needs to exceed the provider's query timeout to pass validation
|
||||
local.WithHostHealthCheckDeadline(90*time.Second),
|
||||
providerOpt,
|
||||
)
|
||||
}
|
||||
|
||||
func (o *NewActorsOpts) getProvider() (local.HostOption, error) {
|
||||
switch {
|
||||
case o.Postgres != nil && o.SQLite != nil:
|
||||
|
||||
@@ -17,7 +17,7 @@ import (
|
||||
func init() {
|
||||
registerTestControllers = []func(apiGroup *gin.RouterGroup, db *gorm.DB, svc *services){
|
||||
func(apiGroup *gin.RouterGroup, db *gorm.DB, svc *services) {
|
||||
testService, err := service.NewTestService(db, svc.appConfigService, svc.jwtService, svc.ldapService, svc.appLockService, svc.fileStorage)
|
||||
testService, err := service.NewTestService(db, svc.actors, svc.appConfigService, svc.jwtService, svc.ldapService, svc.appLockService, svc.fileStorage)
|
||||
if err != nil {
|
||||
slog.Error("Failed to initialize test service", slog.Any("error", err))
|
||||
os.Exit(1)
|
||||
|
||||
@@ -43,6 +43,7 @@ type services struct {
|
||||
webauthnModule *webauthn.Module
|
||||
userSignUpModule *usersignup.Module
|
||||
apiModule *api.Module
|
||||
actors *local.Host
|
||||
}
|
||||
|
||||
// Initializes all services
|
||||
@@ -56,7 +57,9 @@ func initServices(
|
||||
fileStorage storage.FileStorage,
|
||||
scheduler *job.Scheduler,
|
||||
) (svc *services, err error) {
|
||||
svc = &services{}
|
||||
svc = &services{
|
||||
actors: actors,
|
||||
}
|
||||
|
||||
// Init the app config service
|
||||
svc.appConfigService, err = appconfig.NewService(ctx, actors, db)
|
||||
@@ -132,14 +135,22 @@ func initServices(
|
||||
return nil, fmt.Errorf("failed to create API key module: %w", err)
|
||||
}
|
||||
|
||||
svc.userSignUpModule = usersignup.New(usersignup.Dependencies{
|
||||
svc.userSignUpModule, err = usersignup.New(ctx, usersignup.Dependencies{
|
||||
DB: db,
|
||||
Actors: actors,
|
||||
Signer: svc.jwtService,
|
||||
AuditLog: svc.auditLogService,
|
||||
UserCreator: svc.userService,
|
||||
AppConfig: svc.appConfigService,
|
||||
})
|
||||
svc.oneTimeAccessService = service.NewOneTimeAccessService(db, svc.userService, svc.jwtService, svc.auditLogService, svc.emailService)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create user signup module: %w", err)
|
||||
}
|
||||
|
||||
svc.oneTimeAccessService, err = service.NewOneTimeAccessService(actors, db, svc.userService, svc.jwtService, svc.auditLogService, svc.emailService)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create one-time access service: %w", err)
|
||||
}
|
||||
|
||||
svc.versionService = service.NewVersionService(httpClient)
|
||||
|
||||
|
||||
@@ -24,57 +24,47 @@ var oneTimeAccessTokenCmd = &cobra.Command{
|
||||
userArg := args[0]
|
||||
|
||||
// Connect to the database
|
||||
db, _, err := bootstrap.NewDatabase(cmd.Context())
|
||||
db, pg, err := bootstrap.NewDatabase(cmd.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Create the access token
|
||||
var oneTimeAccessToken *model.OneTimeAccessToken
|
||||
err = db.Transaction(func(tx *gorm.DB) error {
|
||||
// Load the user to retrieve the user ID
|
||||
var user model.User
|
||||
queryCtx, queryCancel := context.WithTimeout(cmd.Context(), 10*time.Second)
|
||||
defer queryCancel()
|
||||
txErr := tx.
|
||||
WithContext(queryCtx).
|
||||
Where("username = ? OR email = ?", userArg, userArg).
|
||||
First(&user).
|
||||
Error
|
||||
switch {
|
||||
case errors.Is(txErr, gorm.ErrRecordNotFound):
|
||||
return errors.New("user not found")
|
||||
case txErr != nil:
|
||||
return fmt.Errorf("failed to query for user: %w", txErr)
|
||||
case user.ID == "":
|
||||
return errors.New("invalid user loaded: ID is empty")
|
||||
}
|
||||
// Load the user to retrieve the user ID
|
||||
var user model.User
|
||||
queryCtx, queryCancel := context.WithTimeout(cmd.Context(), 10*time.Second)
|
||||
defer queryCancel()
|
||||
err = db.
|
||||
WithContext(queryCtx).
|
||||
Where("username = ? OR email = ?", userArg, userArg).
|
||||
First(&user).
|
||||
Error
|
||||
switch {
|
||||
case errors.Is(err, gorm.ErrRecordNotFound):
|
||||
return errors.New("user not found")
|
||||
case err != nil:
|
||||
return fmt.Errorf("failed to query for user: %w", err)
|
||||
case user.ID == "":
|
||||
return errors.New("invalid user loaded: ID is empty")
|
||||
}
|
||||
|
||||
// Create a new access token that expires in 1 hour
|
||||
oneTimeAccessToken, txErr = service.NewOneTimeAccessToken(user.ID, time.Hour, false)
|
||||
if txErr != nil {
|
||||
return fmt.Errorf("failed to generate access token: %w", txErr)
|
||||
}
|
||||
|
||||
queryCtx, queryCancel = context.WithTimeout(cmd.Context(), 10*time.Second)
|
||||
defer queryCancel()
|
||||
txErr = tx.
|
||||
WithContext(queryCtx).
|
||||
Create(oneTimeAccessToken).
|
||||
Error
|
||||
if txErr != nil {
|
||||
return fmt.Errorf("failed to save access token: %w", txErr)
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
// One-time access tokens are stored in the actor state store
|
||||
// The CLI doesn't run the full actor host, so it uses a minimal state store to persist the token directly
|
||||
actorStore, err := bootstrap.NewActorStateStore(db, pg)
|
||||
if err != nil {
|
||||
return err
|
||||
return fmt.Errorf("failed to initialize the actor state store: %w", err)
|
||||
}
|
||||
|
||||
// Create a new access token that expires in 1 hour
|
||||
tokenCtx, tokenCancel := context.WithTimeout(cmd.Context(), 10*time.Second)
|
||||
defer tokenCancel()
|
||||
token, _, err := service.StoreOneTimeAccessToken(tokenCtx, actorStore, user.ID, time.Hour, false)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create access token: %w", err)
|
||||
}
|
||||
|
||||
// Print the result
|
||||
fmt.Printf(`A one-time access token valid for 1 hour has been created for "%s".`+"\n", userArg)
|
||||
fmt.Printf("Use the following URL to sign in once: %s/lc/%s\n", common.EnvConfig.AppURL, oneTimeAccessToken.Token)
|
||||
fmt.Printf("Use the following URL to sign in once: %s/lc/%s\n", common.EnvConfig.AppURL, token)
|
||||
|
||||
return nil
|
||||
},
|
||||
|
||||
@@ -15,7 +15,6 @@ import (
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/oidc"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/usersignup"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/webauthn"
|
||||
)
|
||||
|
||||
@@ -34,8 +33,6 @@ func (s *Scheduler) RegisterDbCleanupJobs(ctx context.Context, db *gorm.DB) erro
|
||||
// Use exponential backoff for each DB cleanup job so transient query failures are retried automatically rather than causing an immediate job failure
|
||||
return errors.Join(
|
||||
s.RegisterJob(ctx, "ClearWebauthnSessions", jobDefWithJitter(24*time.Hour), jobs.clearWebauthnSessions, service.RegisterJobOpts{RunImmediately: true, BackOff: newBackOff()}),
|
||||
s.RegisterJob(ctx, "ClearOneTimeAccessTokens", jobDefWithJitter(24*time.Hour), jobs.clearOneTimeAccessTokens, service.RegisterJobOpts{RunImmediately: true, BackOff: newBackOff()}),
|
||||
s.RegisterJob(ctx, "ClearSignupTokens", jobDefWithJitter(24*time.Hour), jobs.clearSignupTokens, service.RegisterJobOpts{RunImmediately: true, BackOff: newBackOff()}),
|
||||
s.RegisterJob(ctx, "ClearEmailVerificationTokens", jobDefWithJitter(24*time.Hour), jobs.clearEmailVerificationTokens, service.RegisterJobOpts{RunImmediately: true, BackOff: newBackOff()}),
|
||||
s.RegisterJob(ctx, "ClearOAuth2Sessions", jobDefWithJitter(24*time.Hour), jobs.clearOAuth2Sessions, service.RegisterJobOpts{RunImmediately: true, BackOff: newBackOff()}),
|
||||
s.RegisterJob(ctx, "ClearOAuth2JTIs", jobDefWithJitter(24*time.Hour), jobs.clearOAuth2JTIs, service.RegisterJobOpts{RunImmediately: true, BackOff: newBackOff()}),
|
||||
@@ -61,32 +58,6 @@ func (j *DbCleanupJobs) clearWebauthnSessions(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClearOneTimeAccessTokens deletes one-time access tokens that have expired
|
||||
func (j *DbCleanupJobs) clearOneTimeAccessTokens(ctx context.Context) error {
|
||||
st := j.db.
|
||||
WithContext(ctx).
|
||||
Delete(&model.OneTimeAccessToken{}, "expires_at < ?", datatype.DateTime(time.Now()))
|
||||
if st.Error != nil {
|
||||
return fmt.Errorf("failed to clean expired one-time access tokens: %w", st.Error)
|
||||
}
|
||||
|
||||
slog.InfoContext(ctx, "Cleaned expired one-time access tokens", slog.Int64("count", st.RowsAffected))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// clearSignupTokens deletes signup tokens that have expired
|
||||
func (j *DbCleanupJobs) clearSignupTokens(ctx context.Context) error {
|
||||
count, err := usersignup.CleanupExpiredSignupTokens(ctx, j.db)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to clean expired signup tokens: %w", err)
|
||||
}
|
||||
|
||||
slog.InfoContext(ctx, "Cleaned expired signup tokens", slog.Int64("count", count))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// clearOAuth2Sessions deletes expired and invalidated OAuth2 sessions.
|
||||
func (j *DbCleanupJobs) clearOAuth2Sessions(ctx context.Context) error {
|
||||
count, err := oidc.CleanupExpiredOAuth2Sessions(ctx, j.db)
|
||||
|
||||
@@ -1,13 +0,0 @@
|
||||
package model
|
||||
|
||||
import datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
|
||||
type OneTimeAccessToken struct {
|
||||
Base
|
||||
Token string
|
||||
DeviceToken *string
|
||||
ExpiresAt datatype.DateTime
|
||||
|
||||
UserID string
|
||||
User User
|
||||
}
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
|
||||
"github.com/go-webauthn/webauthn/protocol"
|
||||
"github.com/google/uuid"
|
||||
"github.com/italypaleale/francis/host/local"
|
||||
"github.com/lestrrat-go/jwx/v3/jwa"
|
||||
"github.com/lestrrat-go/jwx/v3/jwk"
|
||||
"github.com/lestrrat-go/jwx/v3/jwt"
|
||||
@@ -41,6 +42,7 @@ import (
|
||||
|
||||
type TestService struct {
|
||||
db *gorm.DB
|
||||
actors *local.Host
|
||||
jwtService *JwtService
|
||||
appConfigService *appconfig.AppConfigService
|
||||
ldapService *LdapService
|
||||
@@ -56,9 +58,10 @@ const (
|
||||
e2eRefreshTokenExpiredFixtureToken = "X4vqwtRyCUaq51UafHea4Fsg8Km6CAns6vp3tuX4"
|
||||
)
|
||||
|
||||
func NewTestService(db *gorm.DB, appConfigService *appconfig.AppConfigService, jwtService *JwtService, ldapService *LdapService, appLockService *AppLockService, fileStorage storage.FileStorage) (*TestService, error) {
|
||||
func NewTestService(db *gorm.DB, actors *local.Host, appConfigService *appconfig.AppConfigService, jwtService *JwtService, ldapService *LdapService, appLockService *AppLockService, fileStorage storage.FileStorage) (*TestService, error) {
|
||||
s := &TestService{
|
||||
db: db,
|
||||
actors: actors,
|
||||
appConfigService: appConfigService,
|
||||
jwtService: jwtService,
|
||||
ldapService: ldapService,
|
||||
@@ -136,29 +139,6 @@ func (s *TestService) SeedDatabase(baseURL string) error {
|
||||
}
|
||||
}
|
||||
|
||||
oneTimeAccessTokens := []model.OneTimeAccessToken{{
|
||||
Base: model.Base{
|
||||
ID: "bf877753-4ea4-4c9c-bbbd-e198bb201cb8",
|
||||
},
|
||||
Token: "HPe6k6uiDRRVuAQV",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(1 * time.Hour)),
|
||||
UserID: users[0].ID,
|
||||
},
|
||||
{
|
||||
Base: model.Base{
|
||||
ID: "d3afae24-fe2d-4a98-abec-cf0b8525096a",
|
||||
},
|
||||
Token: "YCGDtftvsvYWiXd0",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(-1 * time.Second)), // expired
|
||||
UserID: users[0].ID,
|
||||
},
|
||||
}
|
||||
for _, token := range oneTimeAccessTokens {
|
||||
if err := tx.Create(&token).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
userGroups := []model.UserGroup{
|
||||
{
|
||||
Base: model.Base{
|
||||
@@ -282,15 +262,6 @@ func (s *TestService) SeedDatabase(baseURL string) error {
|
||||
}
|
||||
}
|
||||
|
||||
accessToken := model.OneTimeAccessToken{
|
||||
Token: "one-time-token",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(1 * time.Hour)),
|
||||
UserID: users[0].ID,
|
||||
}
|
||||
if err := tx.Create(&accessToken).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
userAuthorizedClients := []model.UserAuthorizedOidcClient{
|
||||
{
|
||||
Scope: datatype.StringList{"openid", "profile", "email"},
|
||||
@@ -448,53 +419,6 @@ func (s *TestService) SeedDatabase(baseURL string) error {
|
||||
}
|
||||
}
|
||||
|
||||
signupTokens := []usersignup.SignupToken{
|
||||
{
|
||||
Base: model.Base{
|
||||
ID: "a1b2c3d4-e5f6-7890-abcd-ef1234567890",
|
||||
},
|
||||
Token: "VALID1234567890A",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(24 * time.Hour)),
|
||||
UsageLimit: 1,
|
||||
UsageCount: 0,
|
||||
UserGroups: []model.UserGroup{
|
||||
userGroups[0],
|
||||
},
|
||||
},
|
||||
{
|
||||
Base: model.Base{
|
||||
ID: "dc3c9c96-714e-48eb-926e-2d7c7858e6cf",
|
||||
},
|
||||
Token: "PARTIAL567890ABC",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(7 * 24 * time.Hour)),
|
||||
UsageLimit: 5,
|
||||
UsageCount: 2,
|
||||
},
|
||||
{
|
||||
Base: model.Base{
|
||||
ID: "44de1863-ffa5-4db1-9507-4887cd7a1e3f",
|
||||
},
|
||||
Token: "EXPIRED34567890B",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(-24 * time.Hour)), // Expired
|
||||
UsageLimit: 3,
|
||||
UsageCount: 1,
|
||||
},
|
||||
{
|
||||
Base: model.Base{
|
||||
ID: "f1b1678b-7720-4d8b-8f91-1dbff1e2d02b",
|
||||
},
|
||||
Token: "FULLYUSED567890C",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(24 * time.Hour)),
|
||||
UsageLimit: 1,
|
||||
UsageCount: 1, // Usage limit reached
|
||||
},
|
||||
}
|
||||
for _, token := range signupTokens {
|
||||
if err := tx.Create(&token).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
emailVerificationTokens := []model.EmailVerificationToken{
|
||||
{
|
||||
Base: model.Base{
|
||||
@@ -541,6 +465,81 @@ func (s *TestService) SeedDatabase(baseURL string) error {
|
||||
return err
|
||||
}
|
||||
|
||||
// One-time access tokens and signup tokens live in the actor state store, so they're seeded separately from the DB transaction above.
|
||||
err = s.seedOneTimeAccessTokens(context.Background())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to seed one-time access tokens: %w", err)
|
||||
}
|
||||
|
||||
err = s.seedSignupTokens(context.Background())
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to seed signup tokens: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// seedSignupTokens seeds the signup tokens used by E2E tests into the signup token singleton actor.
|
||||
// The already-expired fixture token is intentionally not seeded, since the actor would purge it right away via its cleanup alarm.
|
||||
func (s *TestService) seedSignupTokens(ctx context.Context) error {
|
||||
now := time.Now().Round(time.Second)
|
||||
seeds := []usersignup.SignupTokenSeed{
|
||||
{
|
||||
ID: "a1b2c3d4-e5f6-7890-abcd-ef1234567890",
|
||||
Token: "VALID1234567890A",
|
||||
ExpiresAt: now.Add(24 * time.Hour),
|
||||
UsageLimit: 1,
|
||||
UsageCount: 0,
|
||||
UserGroupIDs: []string{"c7ae7c01-28a3-4f3c-9572-1ee734ea8368"},
|
||||
CreatedAt: now,
|
||||
},
|
||||
{
|
||||
ID: "dc3c9c96-714e-48eb-926e-2d7c7858e6cf",
|
||||
Token: "PARTIAL567890ABC",
|
||||
ExpiresAt: now.Add(7 * 24 * time.Hour),
|
||||
UsageLimit: 5,
|
||||
UsageCount: 2,
|
||||
CreatedAt: now,
|
||||
},
|
||||
{
|
||||
ID: "f1b1678b-7720-4d8b-8f91-1dbff1e2d02b",
|
||||
Token: "FULLYUSED567890C",
|
||||
ExpiresAt: now.Add(24 * time.Hour),
|
||||
UsageLimit: 1,
|
||||
UsageCount: 1, // Usage limit reached
|
||||
CreatedAt: now,
|
||||
},
|
||||
}
|
||||
|
||||
return usersignup.SeedSignupTokens(ctx, s.actors, seeds)
|
||||
}
|
||||
|
||||
// seedOneTimeAccessTokens seeds the one-time access tokens used by E2E tests into the actor state store.
|
||||
// Expired tokens are intentionally not seeded: with actor-backed storage an expired token is simply one that has no state, which the exchange flow already reports as invalid/expired.
|
||||
func (s *TestService) seedOneTimeAccessTokens(ctx context.Context) error {
|
||||
tokens := []struct {
|
||||
token string
|
||||
ttl time.Duration
|
||||
}{
|
||||
{token: "HPe6k6uiDRRVuAQV", ttl: time.Hour},
|
||||
{token: "one-time-token", ttl: time.Hour},
|
||||
}
|
||||
|
||||
for _, t := range tokens {
|
||||
state := oneTimeAccessTokenState{
|
||||
UserID: e2eRefreshTokenUserID,
|
||||
ExpiresAt: time.Now().Add(t.ttl).Round(time.Second),
|
||||
}
|
||||
// Seed through the actor's "restore" method (which sets the state) rather than writing the
|
||||
// state directly: if an actor for this token is still active from a previous test (for
|
||||
// example, one whose token was already consumed), invoking it refreshes its in-memory cache
|
||||
// too, whereas a direct state write would leave that cache stale.
|
||||
_, err := s.actors.Service().Invoke(ctx, OneTimeAccessTokenActorType, t.token, oneTimeAccessTokenMethodRestore, state)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to seed one-time access token %q: %w", t.token, err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
164
backend/internal/service/one_time_access_actor.go
Normal file
164
backend/internal/service/one_time_access_actor.go
Normal file
@@ -0,0 +1,164 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/italypaleale/francis/actor"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/common"
|
||||
)
|
||||
|
||||
// One-time access 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).
|
||||
|
||||
// OneTimeAccessTokenActorType is the actor type for the one-time access token actor
|
||||
const OneTimeAccessTokenActorType = "OneTimeAccessToken"
|
||||
|
||||
// Methods exposed by the one-time access 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 and deletes the token), and only afterwards performing the remaining work.
|
||||
// On failure, the caller compensates by restoring the token via the "restore" method as best-effort.
|
||||
const (
|
||||
oneTimeAccessTokenMethodConsume = "consume"
|
||||
oneTimeAccessTokenMethodRestore = "restore"
|
||||
)
|
||||
|
||||
// oneTimeAccessConsumeStatus is the outcome of a "consume" invocation.
|
||||
type oneTimeAccessConsumeStatus string
|
||||
|
||||
const (
|
||||
// oneTimeAccessConsumeOK indicates the token was valid and has been consumed
|
||||
oneTimeAccessConsumeOK oneTimeAccessConsumeStatus = "ok"
|
||||
// oneTimeAccessConsumeNotFound indicates the token doesn't exist (or has expired)
|
||||
oneTimeAccessConsumeNotFound oneTimeAccessConsumeStatus = "not_found"
|
||||
// oneTimeAccessConsumeDeviceMismatch indicates the provided device token doesn't match
|
||||
oneTimeAccessConsumeDeviceMismatch oneTimeAccessConsumeStatus = "device_mismatch"
|
||||
)
|
||||
|
||||
// oneTimeAccessTokenState is the persisted state of a one-time access token actor
|
||||
type oneTimeAccessTokenState struct {
|
||||
UserID string
|
||||
DeviceToken *string
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
// oneTimeAccessConsumeRequest is the payload for the "consume" method
|
||||
type oneTimeAccessConsumeRequest struct {
|
||||
DeviceToken string
|
||||
}
|
||||
|
||||
// oneTimeAccessConsumeResponse is the response of the "consume" method
|
||||
type oneTimeAccessConsumeResponse struct {
|
||||
Status oneTimeAccessConsumeStatus
|
||||
// State is included only when Status is "ok", so the caller can restore it if a later step fails
|
||||
State oneTimeAccessTokenState
|
||||
}
|
||||
|
||||
// oneTimeAccessTokenActor is the actor that manages a single one-time access token
|
||||
type oneTimeAccessTokenActor struct {
|
||||
log *slog.Logger
|
||||
client actor.Client[oneTimeAccessTokenState]
|
||||
}
|
||||
|
||||
// NewOneTimeAccessTokenActor allocates a new one-time access token actor
|
||||
// It satisfies actor.Factory
|
||||
func NewOneTimeAccessTokenActor(actorID string, service *actor.Service) actor.Actor {
|
||||
return &oneTimeAccessTokenActor{
|
||||
log: slog.With(
|
||||
slog.String("scope", "actor"),
|
||||
slog.String("actorType", OneTimeAccessTokenActorType),
|
||||
),
|
||||
client: actor.NewActorClient[oneTimeAccessTokenState](OneTimeAccessTokenActorType, actorID, service),
|
||||
}
|
||||
}
|
||||
|
||||
// Invoke implements actor.ActorInvoke
|
||||
func (a *oneTimeAccessTokenActor) Invoke(parentCtx context.Context, method string, data actor.Envelope) (any, error) {
|
||||
switch method {
|
||||
case oneTimeAccessTokenMethodConsume:
|
||||
return a.consume(parentCtx, data)
|
||||
case oneTimeAccessTokenMethodRestore:
|
||||
return nil, a.restore(parentCtx, data)
|
||||
default:
|
||||
return nil, common.ErrUnsupportedActorMethod{Method: method}
|
||||
}
|
||||
}
|
||||
|
||||
// consume atomically validates the token and, if valid, deletes it.
|
||||
func (a *oneTimeAccessTokenActor) consume(parentCtx context.Context, data actor.Envelope) (oneTimeAccessConsumeResponse, error) {
|
||||
var req oneTimeAccessConsumeRequest
|
||||
if data != nil {
|
||||
err := data.Decode(&req)
|
||||
if err != nil {
|
||||
return oneTimeAccessConsumeResponse{}, fmt.Errorf("request body is not valid for method '%s': %w", oneTimeAccessTokenMethodConsume, err)
|
||||
}
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
state, err := a.client.GetState(ctx)
|
||||
if err != nil {
|
||||
return oneTimeAccessConsumeResponse{}, fmt.Errorf("error retrieving actor state: %w", err)
|
||||
}
|
||||
|
||||
// An empty UserID means there's no state: the token doesn't exist (or its state already expired and was purged)
|
||||
if state.UserID == "" || state.ExpiresAt.Before(time.Now()) {
|
||||
return oneTimeAccessConsumeResponse{
|
||||
Status: oneTimeAccessConsumeNotFound,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// If the token requires a device token, it must match
|
||||
// A mismatch leaves the token untouched, mirroring the pre-actor behavior
|
||||
if state.DeviceToken != nil && req.DeviceToken != *state.DeviceToken {
|
||||
return oneTimeAccessConsumeResponse{
|
||||
Status: oneTimeAccessConsumeDeviceMismatch,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// The token is valid: delete the state (one-time use)
|
||||
ctx, cancel = context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
err = a.client.DeleteState(ctx)
|
||||
if err != nil {
|
||||
return oneTimeAccessConsumeResponse{}, fmt.Errorf("error deleting actor state: %w", err)
|
||||
}
|
||||
|
||||
return oneTimeAccessConsumeResponse{
|
||||
Status: oneTimeAccessConsumeOK,
|
||||
State: state,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// restore re-creates the token state, used to compensate when a step after consuming the token fails.
|
||||
func (a *oneTimeAccessTokenActor) restore(parentCtx context.Context, data actor.Envelope) error {
|
||||
if data == nil {
|
||||
return fmt.Errorf("request body is empty for method '%s'", oneTimeAccessTokenMethodRestore)
|
||||
}
|
||||
|
||||
var state oneTimeAccessTokenState
|
||||
err := data.Decode(&state)
|
||||
if err != nil {
|
||||
return fmt.Errorf("request body is not valid for method '%s': %w", oneTimeAccessTokenMethodRestore, err)
|
||||
}
|
||||
|
||||
// If the token has meanwhile expired, there's nothing to restore
|
||||
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
|
||||
}
|
||||
@@ -9,32 +9,46 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/italypaleale/francis/actor"
|
||||
"github.com/italypaleale/francis/host/local"
|
||||
"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/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils/email"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/clause"
|
||||
)
|
||||
|
||||
// OneTimeAccessTokenStore is the minimal interface needed to persist a one-time access token in the actor state store.
|
||||
// It's satisfied by both *actor.Service (used by the running application) and *local.Host (used by CLI commands, which don't run the full actor host).
|
||||
type OneTimeAccessTokenStore interface {
|
||||
SetState(ctx context.Context, actorType string, actorID string, state any, opts *actor.SetStateOpts) error
|
||||
}
|
||||
|
||||
type OneTimeAccessService struct {
|
||||
db *gorm.DB
|
||||
actorService *actor.Service
|
||||
userService *UserService
|
||||
jwtService *JwtService
|
||||
auditLogService *AuditLogService
|
||||
emailService *EmailService
|
||||
}
|
||||
|
||||
func NewOneTimeAccessService(db *gorm.DB, userService *UserService, jwtService *JwtService, auditLogService *AuditLogService, emailService *EmailService) *OneTimeAccessService {
|
||||
func NewOneTimeAccessService(actors *local.Host, db *gorm.DB, userService *UserService, jwtService *JwtService, auditLogService *AuditLogService, emailService *EmailService) (*OneTimeAccessService, error) {
|
||||
err := actors.RegisterActor(OneTimeAccessTokenActorType, NewOneTimeAccessTokenActor)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error registering the %s actor: %w", OneTimeAccessTokenActorType, err)
|
||||
}
|
||||
|
||||
return &OneTimeAccessService{
|
||||
db: db,
|
||||
actorService: actors.Service(),
|
||||
userService: userService,
|
||||
jwtService: jwtService,
|
||||
auditLogService: auditLogService,
|
||||
emailService: emailService,
|
||||
}
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsAdmin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID string, ttl time.Duration) error {
|
||||
@@ -71,12 +85,8 @@ func (s *OneTimeAccessService) RequestOneTimeAccessEmailAsUnauthenticatedUser(ct
|
||||
}
|
||||
|
||||
func (s *OneTimeAccessService) requestOneTimeAccessEmailInternal(ctx context.Context, userID, redirectPath string, ttl time.Duration, withDeviceToken bool, dbConfig *appconfig.AppConfigModel) (*string, error) {
|
||||
tx := s.db.Begin()
|
||||
defer func() {
|
||||
tx.Rollback()
|
||||
}()
|
||||
|
||||
user, err := s.userService.getUserInternal(ctx, userID, tx)
|
||||
// Load the user to ensure it exists and has an email address
|
||||
user, err := s.userService.GetUser(ctx, userID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -85,11 +95,7 @@ func (s *OneTimeAccessService) requestOneTimeAccessEmailInternal(ctx context.Con
|
||||
return nil, &common.UserEmailNotSetError{}
|
||||
}
|
||||
|
||||
oneTimeAccessToken, deviceToken, err := s.createOneTimeAccessTokenInternal(ctx, user.ID, ttl, withDeviceToken, tx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
err = tx.Commit().Error
|
||||
oneTimeAccessToken, deviceToken, err := StoreOneTimeAccessToken(ctx, s.actorService, user.ID, ttl, withDeviceToken)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -127,77 +133,79 @@ func (s *OneTimeAccessService) requestOneTimeAccessEmailInternal(ctx context.Con
|
||||
}
|
||||
|
||||
func (s *OneTimeAccessService) CreateOneTimeAccessToken(ctx context.Context, userID string, ttl time.Duration) (token string, err error) {
|
||||
tx := s.db.Begin()
|
||||
defer func() {
|
||||
tx.Rollback()
|
||||
}()
|
||||
|
||||
// Load the user to ensure it exists
|
||||
_, err = s.userService.getUserInternal(ctx, userID, tx)
|
||||
_, err = s.userService.GetUser(ctx, userID)
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return "", &common.UserNotFoundError{}
|
||||
} else if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Create the one-time access token
|
||||
token, _, err = s.createOneTimeAccessTokenInternal(ctx, userID, ttl, false, tx)
|
||||
token, _, err = StoreOneTimeAccessToken(ctx, s.actorService, userID, ttl, false)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Commit
|
||||
err = tx.Commit().Error
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("error committing transaction: %w", err)
|
||||
}
|
||||
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func (s *OneTimeAccessService) createOneTimeAccessTokenInternal(ctx context.Context, userID string, ttl time.Duration, withDeviceToken bool, tx *gorm.DB) (token string, deviceToken *string, err error) {
|
||||
oneTimeAccessToken, err := NewOneTimeAccessToken(userID, ttl, withDeviceToken)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
err = tx.WithContext(ctx).Create(oneTimeAccessToken).Error
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
return oneTimeAccessToken.Token, oneTimeAccessToken.DeviceToken, nil
|
||||
}
|
||||
|
||||
func (s *OneTimeAccessService) ExchangeOneTimeAccessToken(ctx context.Context, dbConfig *appconfig.AppConfigModel, token, deviceToken, ipAddress, userAgent string) (model.User, string, error) {
|
||||
tx := s.db.Begin()
|
||||
defer func() {
|
||||
tx.Rollback()
|
||||
}()
|
||||
|
||||
var oneTimeAccessToken model.OneTimeAccessToken
|
||||
err := tx.
|
||||
WithContext(ctx).
|
||||
Where("token = ? AND expires_at > ?", token, datatype.DateTime(time.Now())).
|
||||
Preload("User").
|
||||
Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
First(&oneTimeAccessToken).
|
||||
Error
|
||||
// Consume the token by invoking its actor: this atomically validates it and, if valid, deletes it.
|
||||
// It must happen outside of a DB transaction, since invoking an actor while a transaction is open would deadlock on SQLite.
|
||||
res, err := s.actorService.Invoke(ctx, OneTimeAccessTokenActorType, token, oneTimeAccessTokenMethodConsume, oneTimeAccessConsumeRequest{
|
||||
DeviceToken: deviceToken,
|
||||
})
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return model.User{}, "", &common.TokenInvalidOrExpiredError{}
|
||||
}
|
||||
return model.User{}, "", fmt.Errorf("error invoking one-time access token actor: %w", err)
|
||||
}
|
||||
|
||||
var consumeRes oneTimeAccessConsumeResponse
|
||||
err = res.Decode(&consumeRes)
|
||||
if err != nil {
|
||||
return model.User{}, "", fmt.Errorf("error decoding one-time access token actor response: %w", err)
|
||||
}
|
||||
|
||||
switch consumeRes.Status {
|
||||
case oneTimeAccessConsumeNotFound:
|
||||
return model.User{}, "", &common.TokenInvalidOrExpiredError{}
|
||||
case oneTimeAccessConsumeDeviceMismatch:
|
||||
return model.User{}, "", &common.DeviceCodeInvalid{}
|
||||
case oneTimeAccessConsumeOK:
|
||||
// All good, continue below
|
||||
default:
|
||||
return model.User{}, "", fmt.Errorf("unexpected status from one-time access token actor: %s", consumeRes.Status)
|
||||
}
|
||||
|
||||
// The token has now been consumed. From this point on, if we hit an error we compensate by restoring the token (this is best-effort).
|
||||
user, accessToken, err := s.completeOneTimeAccessTokenExchange(ctx, dbConfig, consumeRes.State, ipAddress, userAgent)
|
||||
if err != nil {
|
||||
s.restoreOneTimeAccessToken(ctx, token, consumeRes.State)
|
||||
return model.User{}, "", err
|
||||
}
|
||||
if oneTimeAccessToken.DeviceToken != nil && deviceToken != *oneTimeAccessToken.DeviceToken {
|
||||
return model.User{}, "", &common.DeviceCodeInvalid{}
|
||||
|
||||
return user, accessToken, nil
|
||||
}
|
||||
|
||||
// completeOneTimeAccessTokenExchange performs the work that follows consuming a token: loading the user, validating it, and issuing an access token.
|
||||
func (s *OneTimeAccessService) completeOneTimeAccessTokenExchange(ctx context.Context, dbConfig *appconfig.AppConfigModel, state oneTimeAccessTokenState, ipAddress, userAgent string) (model.User, string, error) {
|
||||
var user model.User
|
||||
err := s.db.
|
||||
WithContext(ctx).
|
||||
Where("id = ?", state.UserID).
|
||||
First(&user).
|
||||
Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return model.User{}, "", &common.TokenInvalidOrExpiredError{}
|
||||
} else if err != nil {
|
||||
return model.User{}, "", err
|
||||
}
|
||||
if oneTimeAccessToken.User.Disabled {
|
||||
|
||||
if user.Disabled {
|
||||
return model.User{}, "", &common.UserDisabledError{}
|
||||
}
|
||||
|
||||
accessToken, err := s.jwtService.GenerateAccessToken(
|
||||
oneTimeAccessToken.User,
|
||||
user,
|
||||
AuthenticationMethodOneTimePassword,
|
||||
dbConfig.SessionDuration.AsDurationMinutes(),
|
||||
)
|
||||
@@ -205,57 +213,73 @@ func (s *OneTimeAccessService) ExchangeOneTimeAccessToken(ctx context.Context, d
|
||||
return model.User{}, "", err
|
||||
}
|
||||
|
||||
err = tx.
|
||||
WithContext(ctx).
|
||||
Delete(&oneTimeAccessToken).
|
||||
Error
|
||||
if err != nil {
|
||||
return model.User{}, "", err
|
||||
}
|
||||
|
||||
s.auditLogService.Create(
|
||||
ctx, model.AuditLogEventOneTimeAccessTokenSignIn,
|
||||
ipAddress, userAgent,
|
||||
oneTimeAccessToken.User.ID, model.AuditLogData{},
|
||||
tx,
|
||||
user.ID,
|
||||
model.AuditLogData{},
|
||||
s.db,
|
||||
)
|
||||
|
||||
err = tx.Commit().Error
|
||||
if err != nil {
|
||||
return model.User{}, "", fmt.Errorf("error committing transaction: %w", err)
|
||||
}
|
||||
|
||||
return oneTimeAccessToken.User, accessToken, nil
|
||||
return user, accessToken, nil
|
||||
}
|
||||
|
||||
func NewOneTimeAccessToken(userID string, ttl time.Duration, withDeviceToken bool) (*model.OneTimeAccessToken, error) {
|
||||
// restoreOneTimeAccessToken restores a token that was consumed but whose exchange could not be completed.
|
||||
// It's a best-effort compensation: if it fails (or the process crashes before it runs) we accept that the token was consumed unnecessarily.
|
||||
func (s *OneTimeAccessService) restoreOneTimeAccessToken(parentCtx context.Context, token string, state oneTimeAccessTokenState) {
|
||||
// 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, OneTimeAccessTokenActorType, token, oneTimeAccessTokenMethodRestore, state)
|
||||
if err != nil {
|
||||
slog.ErrorContext(ctx, "Failed to restore one-time access token after a failed exchange", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
|
||||
// StoreOneTimeAccessToken generates a new one-time access token and persists it in the actor state store, with a TTL matching its lifetime.
|
||||
// It returns the token value and, when requested, the associated device token.
|
||||
func StoreOneTimeAccessToken(ctx context.Context, store OneTimeAccessTokenStore, userID string, ttl time.Duration, withDeviceToken bool) (token string, deviceToken *string, err error) {
|
||||
token, deviceToken, err = generateOneTimeAccessToken(ttl, withDeviceToken)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
now := time.Now().Round(time.Second)
|
||||
state := oneTimeAccessTokenState{
|
||||
UserID: userID,
|
||||
DeviceToken: deviceToken,
|
||||
ExpiresAt: now.Add(ttl),
|
||||
}
|
||||
|
||||
err = store.SetState(ctx, OneTimeAccessTokenActorType, token, state, &actor.SetStateOpts{TTL: ttl})
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("error saving one-time access token state: %w", err)
|
||||
}
|
||||
|
||||
return token, deviceToken, nil
|
||||
}
|
||||
|
||||
// generateOneTimeAccessToken generates the random token value (and optional device token) for a one-time access token.
|
||||
func generateOneTimeAccessToken(ttl time.Duration, withDeviceToken bool) (token string, deviceToken *string, err error) {
|
||||
// If expires at is less than 15 minutes, use a 6-character token instead of 16
|
||||
tokenLength := 16
|
||||
if ttl <= 15*time.Minute {
|
||||
tokenLength = 6
|
||||
}
|
||||
|
||||
token, err := utils.GenerateRandomUnambiguousString(tokenLength)
|
||||
token, err = utils.GenerateRandomUnambiguousString(tokenLength)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
var deviceToken *string
|
||||
if withDeviceToken {
|
||||
dt, err := utils.GenerateRandomAlphanumericString(16)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return "", nil, err
|
||||
}
|
||||
deviceToken = &dt
|
||||
}
|
||||
|
||||
now := time.Now().Round(time.Second)
|
||||
o := &model.OneTimeAccessToken{
|
||||
UserID: userID,
|
||||
ExpiresAt: datatype.DateTime(now.Add(ttl)),
|
||||
Token: token,
|
||||
DeviceToken: deviceToken,
|
||||
}
|
||||
|
||||
return o, nil
|
||||
return token, deviceToken, nil
|
||||
}
|
||||
|
||||
@@ -4,22 +4,111 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/italypaleale/francis/actor"
|
||||
"github.com/italypaleale/francis/host/local"
|
||||
"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/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
)
|
||||
|
||||
func TestExchangeOneTimeAccessTokenRejectsDisabledUser(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
// newOneTimeAccessServiceForTest sets up a OneTimeAccessService backed by an in-memory test actor host and returns both.
|
||||
func newOneTimeAccessServiceForTest(t *testing.T, db *gorm.DB) (*OneTimeAccessService, *local.Host) {
|
||||
t.Helper()
|
||||
|
||||
appConfig := appconfig.NewTestAppConfigService(nil)
|
||||
instanceID := newInstanceID(t, db)
|
||||
jwtService := initJwtService(t, db, instanceID, appConfig, newTestEnvConfig())
|
||||
auditLogService := NewAuditLogService(db, nil, &GeoLiteService{}, appConfig)
|
||||
oneTimeAccessService := NewOneTimeAccessService(db, nil, jwtService, auditLogService, nil)
|
||||
|
||||
var svc *OneTimeAccessService
|
||||
host := testutils.NewActorHostForTest(t, func(t *testing.T, h *local.Host) {
|
||||
var err error
|
||||
svc, err = NewOneTimeAccessService(h, db, nil, jwtService, auditLogService, nil)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
require.NotNil(t, svc)
|
||||
|
||||
return svc, host
|
||||
}
|
||||
|
||||
func TestExchangeOneTimeAccessTokenSuccess(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
oneTimeAccessService, host := newOneTimeAccessServiceForTest(t, db)
|
||||
|
||||
user := model.User{
|
||||
Base: model.Base{ID: "enabled-user"},
|
||||
Username: "enabled-user",
|
||||
}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
|
||||
token, _, err := StoreOneTimeAccessToken(t.Context(), oneTimeAccessService.actorService, user.ID, time.Minute, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
exchangedUser, accessToken, err := oneTimeAccessService.ExchangeOneTimeAccessToken(t.Context(), dbConfig, token, "", "1.2.3.4", "test-agent")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, user.ID, exchangedUser.ID)
|
||||
require.NotEmpty(t, accessToken)
|
||||
|
||||
// The token must have been consumed
|
||||
var state oneTimeAccessTokenState
|
||||
err = host.GetState(t.Context(), OneTimeAccessTokenActorType, token, &state)
|
||||
require.ErrorIs(t, err, actor.ErrStateNotFound)
|
||||
|
||||
// A sign-in audit log must have been created
|
||||
var auditLogCount int64
|
||||
require.NoError(t, db.Model(&model.AuditLog{}).
|
||||
Where("user_id = ? AND event = ?", user.ID, model.AuditLogEventOneTimeAccessTokenSignIn).
|
||||
Count(&auditLogCount).Error)
|
||||
require.Equal(t, int64(1), auditLogCount)
|
||||
}
|
||||
|
||||
func TestExchangeOneTimeAccessTokenInvalidToken(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
oneTimeAccessService, _ := newOneTimeAccessServiceForTest(t, db)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
_, _, err := oneTimeAccessService.ExchangeOneTimeAccessToken(t.Context(), dbConfig, "does-not-exist", "", "", "")
|
||||
|
||||
var invalidErr *common.TokenInvalidOrExpiredError
|
||||
require.ErrorAs(t, err, &invalidErr)
|
||||
}
|
||||
|
||||
func TestExchangeOneTimeAccessTokenDeviceMismatch(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
oneTimeAccessService, host := newOneTimeAccessServiceForTest(t, db)
|
||||
|
||||
user := model.User{
|
||||
Base: model.Base{ID: "device-user"},
|
||||
Username: "device-user",
|
||||
}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
|
||||
// Store a token that requires a device token
|
||||
token, deviceToken, err := StoreOneTimeAccessToken(t.Context(), oneTimeAccessService.actorService, user.ID, time.Minute, true)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, deviceToken)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
_, _, err = oneTimeAccessService.ExchangeOneTimeAccessToken(t.Context(), dbConfig, token, "wrong-device-token", "", "")
|
||||
|
||||
var deviceErr *common.DeviceCodeInvalid
|
||||
require.ErrorAs(t, err, &deviceErr)
|
||||
|
||||
// The token must not have been consumed on a device-token mismatch
|
||||
var state oneTimeAccessTokenState
|
||||
err = host.GetState(t.Context(), OneTimeAccessTokenActorType, token, &state)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, user.ID, state.UserID)
|
||||
}
|
||||
|
||||
func TestExchangeOneTimeAccessTokenRejectsDisabledUser(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
oneTimeAccessService, host := newOneTimeAccessServiceForTest(t, db)
|
||||
|
||||
user := model.User{
|
||||
Base: model.Base{ID: "disabled-user"},
|
||||
@@ -28,24 +117,23 @@ func TestExchangeOneTimeAccessTokenRejectsDisabledUser(t *testing.T) {
|
||||
}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
|
||||
loginCode := model.OneTimeAccessToken{
|
||||
Base: model.Base{ID: "disabled-user-login-code"},
|
||||
Token: "ABCDEF",
|
||||
ExpiresAt: datatype.DateTime(time.Now().Add(time.Minute)),
|
||||
UserID: user.ID,
|
||||
}
|
||||
require.NoError(t, db.Create(&loginCode).Error)
|
||||
// Store a one-time access token for the disabled user in the actor state store
|
||||
token, _, err := StoreOneTimeAccessToken(t.Context(), oneTimeAccessService.actorService, user.ID, time.Minute, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
exchangedUser, accessToken, err := oneTimeAccessService.ExchangeOneTimeAccessToken(t.Context(), dbConfig, loginCode.Token, "", "", "")
|
||||
exchangedUser, accessToken, err := oneTimeAccessService.ExchangeOneTimeAccessToken(t.Context(), dbConfig, token, "", "", "")
|
||||
|
||||
var userDisabledErr *common.UserDisabledError
|
||||
require.ErrorAs(t, err, &userDisabledErr)
|
||||
require.Empty(t, exchangedUser.ID)
|
||||
require.Empty(t, accessToken)
|
||||
|
||||
var remainingLoginCode model.OneTimeAccessToken
|
||||
require.NoError(t, db.Where("token = ?", loginCode.Token).First(&remainingLoginCode).Error)
|
||||
// The token must have been restored (not consumed), since the exchange failed because the user is disabled
|
||||
var state oneTimeAccessTokenState
|
||||
err = host.GetState(t.Context(), OneTimeAccessTokenActorType, token, &state)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, user.ID, state.UserID)
|
||||
|
||||
var auditLogCount int64
|
||||
require.NoError(t, db.Model(&model.AuditLog{}).Where("user_id = ?", user.ID).Count(&auditLogCount).Error)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
74
backend/internal/usersignup/migration.go
Normal file
74
backend/internal/usersignup/migration.go
Normal file
@@ -0,0 +1,74 @@
|
||||
package usersignup
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"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 "move tokens to actor state" 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 seed the singleton signup token actor 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"`
|
||||
}
|
||||
|
||||
// loadMigratedSignupTokens reads the signup tokens frozen into the kv table by the migration, so the singleton actor can seed its state from them on first startup
|
||||
// It returns nil if there's nothing to migrate
|
||||
func loadMigratedSignupTokens(ctx context.Context, db *gorm.DB) ([]storedSignupToken, 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)
|
||||
}
|
||||
|
||||
tokens := make([]storedSignupToken, len(migrated))
|
||||
for i, m := range migrated {
|
||||
tokens[i] = storedSignupToken{
|
||||
ID: m.ID,
|
||||
Token: m.Token,
|
||||
ExpiresAt: time.Unix(m.ExpiresAt, 0),
|
||||
UsageLimit: m.UsageLimit,
|
||||
UsageCount: m.UsageCount,
|
||||
UserGroupIDs: m.UserGroupIDs,
|
||||
CreatedAt: time.Unix(m.CreatedAt, 0),
|
||||
}
|
||||
}
|
||||
|
||||
return tokens, nil
|
||||
}
|
||||
155
backend/internal/usersignup/migration_test.go
Normal file
155
backend/internal/usersignup/migration_test.go
Normal file
@@ -0,0 +1,155 @@
|
||||
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 "move tokens to actor state" migration.
|
||||
const versionBeforeMoveTokens = 20260718000000
|
||||
|
||||
// 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]storedSignupToken, 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.Unix())
|
||||
require.Equal(t, createdAt.Unix(), tok1.CreatedAt.Unix())
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
// 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,31 @@ type Module struct {
|
||||
handler *handler
|
||||
}
|
||||
|
||||
func New(deps Dependencies) *Module {
|
||||
service := newService(deps)
|
||||
func New(ctx context.Context, deps Dependencies) (*Module, error) {
|
||||
// Load the signup tokens frozen into the kv table by the migration, so the singleton actor can be seeded (migrated) from them on first startup
|
||||
migrated, err := loadMigratedSignupTokens(ctx, deps.DB)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Register the singleton actor that holds all signup tokens
|
||||
bootstrapData := &signupTokenBootstrap{
|
||||
Tokens: migrated,
|
||||
}
|
||||
err = deps.Actors.RegisterSingletonActor(
|
||||
SignupTokenActorType, NewSignupTokenActor,
|
||||
local.WithBootstrapData(bootstrapData),
|
||||
local.WithIdleTimeout(-1), // Disable idle timeout for this actor
|
||||
)
|
||||
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
|
||||
}
|
||||
|
||||
// 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,51 @@ 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, actor.SingletonActorID, signupTokenMethodConsume, signupTokenConsumeRequest{
|
||||
Token: signupData.Token,
|
||||
})
|
||||
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 +82,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 +115,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 +131,21 @@ 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, actor.SingletonActorID, signupTokenMethodRelease, signupTokenReleaseRequest{
|
||||
Token: token,
|
||||
})
|
||||
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 +203,217 @@ 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{})
|
||||
// Signup tokens are held in the singleton actor's state, so we retrieve them all via a read-only Peek and then sort and paginate in memory.
|
||||
res, err := s.actorService.Peek(ctx, SignupTokenActorType, actor.SingletonActorID, signupTokenMethodList, nil)
|
||||
if err != nil {
|
||||
return nil, utils.PaginationResponse{}, fmt.Errorf("error listing signup tokens from actor: %w", err)
|
||||
}
|
||||
|
||||
pagination, err := utils.PaginateFilterAndSort(listRequestOptions, query, &tokens)
|
||||
return tokens, pagination, err
|
||||
var listRes signupTokenListResponse
|
||||
err = res.Decode(&listRes)
|
||||
if err != nil {
|
||||
return nil, utils.PaginationResponse{}, fmt.Errorf("error decoding signup token actor response: %w", err)
|
||||
}
|
||||
|
||||
// Resolve the referenced user groups so they can be included in the response
|
||||
groupsByID, err := s.loadUserGroupsByID(ctx, listRes.Tokens)
|
||||
if err != nil {
|
||||
return nil, utils.PaginationResponse{}, err
|
||||
}
|
||||
|
||||
tokens := make([]SignupToken, len(listRes.Tokens))
|
||||
for i, t := range listRes.Tokens {
|
||||
tokens[i] = signupTokenModelFromStored(t, resolveUserGroups(t.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
|
||||
_, err := s.actorService.Invoke(ctx, SignupTokenActorType, actor.SingletonActorID, signupTokenMethodDelete, signupTokenDeleteRequest{
|
||||
ID: tokenID,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("error deleting signup token via actor: %w", err)
|
||||
}
|
||||
return 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)
|
||||
stored := storedSignupToken{
|
||||
ID: uuid.NewString(),
|
||||
Token: randomString,
|
||||
ExpiresAt: now.Add(ttl),
|
||||
UsageLimit: usageLimit,
|
||||
UsageCount: 0,
|
||||
UserGroupIDs: validGroupIDs,
|
||||
CreatedAt: now,
|
||||
}
|
||||
|
||||
_, err = s.actorService.Invoke(ctx, SignupTokenActorType, actor.SingletonActorID, signupTokenMethodCreate, stored)
|
||||
if err != nil {
|
||||
return SignupToken{}, fmt.Errorf("error creating signup token via actor: %w", err)
|
||||
}
|
||||
|
||||
return signupTokenModelFromStored(stored, userGroups), nil
|
||||
}
|
||||
|
||||
// loadUserGroupsByID loads every user group referenced by the given tokens, keyed by ID.
|
||||
func (s *Service) loadUserGroupsByID(ctx context.Context, tokens []storedSignupToken) (map[string]model.UserGroup, error) {
|
||||
idSet := make(map[string]struct{})
|
||||
for _, t := range tokens {
|
||||
for _, id := range t.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
|
||||
}
|
||||
|
||||
// signupTokenModelFromStored builds the API/model representation of a signup token from its stored form.
|
||||
func signupTokenModelFromStored(t storedSignupToken, groups []model.UserGroup) SignupToken {
|
||||
return SignupToken{
|
||||
Base: model.Base{
|
||||
ID: t.ID,
|
||||
CreatedAt: datatype.DateTime(t.CreatedAt),
|
||||
},
|
||||
Token: t.Token,
|
||||
ExpiresAt: datatype.DateTime(t.ExpiresAt),
|
||||
UsageLimit: t.UsageLimit,
|
||||
UsageCount: t.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
backend/internal/usersignup/service_test.go
Normal file
127
backend/internal/usersignup/service_test.go
Normal 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, nil)
|
||||
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
|
||||
}
|
||||
496
backend/internal/usersignup/signuptoken_actor.go
Normal file
496
backend/internal/usersignup/signuptoken_actor.go
Normal file
@@ -0,0 +1,496 @@
|
||||
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 in a single singleton actor that holds all of them in its state.
|
||||
// The actor keeps a single alarm scheduled for the moment the earliest-expiring token expires. When that alarm fires, the actor purges every expired token and, if any tokens remain, reschedules the alarm for the next earliest expiration.
|
||||
// Because it's a singleton, read-only operations (such as listing tokens) are served via Peek, while mutations (create, delete, consume, release) go through Invoke.
|
||||
//
|
||||
// Consuming a token must happen outside of a DB transaction (invoking an actor while a transaction is open would deadlock on SQLite): the caller invokes the actor to atomically increment the usage count, performs the rest of its work in a transaction, and, on failure, compensates by invoking the actor again to decrement the usage count (best-effort).
|
||||
|
||||
// SignupTokenActorType is the actor type for the signup token singleton actor
|
||||
const SignupTokenActorType = "SignupToken"
|
||||
|
||||
// cleanupAlarmName is the name of the alarm used to purge expired tokens
|
||||
const cleanupAlarmName = "cleanup"
|
||||
|
||||
// Methods exposed by the signup token actor
|
||||
const (
|
||||
signupTokenMethodCreate = "create"
|
||||
signupTokenMethodDelete = "delete"
|
||||
signupTokenMethodConsume = "consume"
|
||||
signupTokenMethodRelease = "release"
|
||||
signupTokenMethodReplace = "replace"
|
||||
signupTokenMethodList = "list"
|
||||
)
|
||||
|
||||
// signupTokenConsumeStatus is the outcome of a "consume" invocation.
|
||||
// It's returned as part of the response payload rather than as a Go error, because errors lose their concrete type when they cross the actor invocation boundary.
|
||||
type signupTokenConsumeStatus string
|
||||
|
||||
const (
|
||||
signupTokenConsumeOK signupTokenConsumeStatus = "ok"
|
||||
signupTokenConsumeNotFound signupTokenConsumeStatus = "not_found"
|
||||
signupTokenConsumeExpired signupTokenConsumeStatus = "expired"
|
||||
signupTokenConsumeLimitReached signupTokenConsumeStatus = "limit_reached"
|
||||
)
|
||||
|
||||
// storedSignupToken is a single signup token as held in the actor state
|
||||
type storedSignupToken struct {
|
||||
ID string
|
||||
Token string
|
||||
ExpiresAt time.Time
|
||||
UsageLimit int
|
||||
UsageCount int
|
||||
UserGroupIDs []string
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
func (t storedSignupToken) isExpired(now time.Time) bool {
|
||||
return t.ExpiresAt.Before(now)
|
||||
}
|
||||
|
||||
func (t storedSignupToken) isUsageLimitReached() bool {
|
||||
return t.UsageCount >= t.UsageLimit
|
||||
}
|
||||
|
||||
// signupTokenActorState is the persisted state of the signup token singleton actor.
|
||||
// Tokens are keyed by their token value.
|
||||
type signupTokenActorState struct {
|
||||
Tokens map[string]storedSignupToken
|
||||
}
|
||||
|
||||
// removeExpired deletes every token that has expired and returns the number removed.
|
||||
func (s *signupTokenActorState) removeExpired(now time.Time) (removed int) {
|
||||
for k, t := range s.Tokens {
|
||||
if t.isExpired(now) {
|
||||
delete(s.Tokens, k)
|
||||
removed++
|
||||
}
|
||||
}
|
||||
|
||||
return removed
|
||||
}
|
||||
|
||||
// earliestExpiration returns the earliest expiration time among all tokens, and whether there's at least one token.
|
||||
func (s *signupTokenActorState) earliestExpiration() (earliest time.Time, found bool) {
|
||||
for _, t := range s.Tokens {
|
||||
if !found || t.ExpiresAt.Before(earliest) {
|
||||
earliest = t.ExpiresAt
|
||||
found = true
|
||||
}
|
||||
}
|
||||
|
||||
return earliest, found
|
||||
}
|
||||
|
||||
// Payloads for the actor methods
|
||||
|
||||
type signupTokenBootstrap struct {
|
||||
Tokens []storedSignupToken
|
||||
}
|
||||
|
||||
type signupTokenDeleteRequest struct {
|
||||
ID string
|
||||
}
|
||||
|
||||
type signupTokenConsumeRequest struct {
|
||||
Token string
|
||||
}
|
||||
|
||||
type signupTokenConsumeResponse struct {
|
||||
Status signupTokenConsumeStatus
|
||||
UserGroupIDs []string
|
||||
}
|
||||
|
||||
type signupTokenReleaseRequest struct {
|
||||
Token string
|
||||
}
|
||||
|
||||
type signupTokenReplaceRequest struct {
|
||||
Tokens []storedSignupToken
|
||||
}
|
||||
|
||||
type signupTokenListResponse struct {
|
||||
Tokens []storedSignupToken
|
||||
}
|
||||
|
||||
// signupTokenActor is the singleton actor that manages all signup tokens
|
||||
type signupTokenActor struct {
|
||||
log *slog.Logger
|
||||
client actor.Client[*signupTokenActorState]
|
||||
}
|
||||
|
||||
// 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),
|
||||
slog.String("actorID", actorID),
|
||||
),
|
||||
client: actor.NewActorClient[*signupTokenActorState](SignupTokenActorType, actorID, service),
|
||||
}
|
||||
}
|
||||
|
||||
// Bootstrap implements actor.ActorBootstrapper for the singleton actor.
|
||||
// On first startup it seeds the state from the tokens migrated from the database
|
||||
// On subsequent startups it just makes sure the cleanup alarm is scheduled.
|
||||
func (a *signupTokenActor) Bootstrap(parentCtx context.Context, data actor.Envelope) 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)
|
||||
}
|
||||
|
||||
// If we already have a state, just make sure the cleanup alarm is scheduled and we're done
|
||||
if state != nil {
|
||||
return a.scheduleCleanup(parentCtx, state)
|
||||
}
|
||||
|
||||
// Initialize the state, seeding it from the migrated tokens (if any)
|
||||
state = &signupTokenActorState{
|
||||
Tokens: map[string]storedSignupToken{},
|
||||
}
|
||||
if data != nil {
|
||||
payload := signupTokenBootstrap{}
|
||||
err = data.Decode(&payload)
|
||||
if err != nil {
|
||||
return fmt.Errorf("request body is not valid for bootstrap: %w", err)
|
||||
}
|
||||
|
||||
for _, t := range payload.Tokens {
|
||||
state.Tokens[t.Token] = t
|
||||
}
|
||||
}
|
||||
|
||||
// Don't carry over tokens that have already expired
|
||||
state.removeExpired(time.Now())
|
||||
|
||||
ctx, cancel = context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
err = a.client.SetState(ctx, state, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("error saving actor state: %w", err)
|
||||
}
|
||||
|
||||
return a.scheduleCleanup(parentCtx, state)
|
||||
}
|
||||
|
||||
// Peek implements actor.ActorPeek for read-only operations
|
||||
func (a *signupTokenActor) Peek(parentCtx context.Context, method string, data actor.Envelope) (any, error) {
|
||||
if method != signupTokenMethodList {
|
||||
return nil, common.ErrUnsupportedActorMethod{Method: method}
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
state, err := a.client.GetState(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error retrieving actor state: %w", err)
|
||||
}
|
||||
|
||||
return signupTokenListResponse{
|
||||
Tokens: collectTokens(state),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Invoke implements actor.ActorInvoke for mutating operations
|
||||
func (a *signupTokenActor) Invoke(parentCtx context.Context, method string, data actor.Envelope) (any, error) {
|
||||
switch method {
|
||||
case signupTokenMethodCreate:
|
||||
return a.create(parentCtx, data)
|
||||
case signupTokenMethodDelete:
|
||||
return nil, a.delete(parentCtx, data)
|
||||
case signupTokenMethodConsume:
|
||||
return a.consume(parentCtx, data)
|
||||
case signupTokenMethodRelease:
|
||||
return nil, a.release(parentCtx, data)
|
||||
case signupTokenMethodReplace:
|
||||
return nil, a.replace(parentCtx, data)
|
||||
default:
|
||||
return nil, common.ErrUnsupportedActorMethod{Method: method}
|
||||
}
|
||||
}
|
||||
|
||||
// Alarm implements actor.ActorAlarm: it purges expired tokens and reschedules the alarm.
|
||||
func (a *signupTokenActor) Alarm(parentCtx context.Context, name string, data actor.Envelope) error {
|
||||
if name != cleanupAlarmName {
|
||||
return common.ErrUnsupportedActorMethod{Method: name}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
if state == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
removed := state.removeExpired(time.Now())
|
||||
if removed > 0 {
|
||||
ctx, cancel = context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
err = a.client.SetState(ctx, state, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("error saving actor state: %w", err)
|
||||
}
|
||||
|
||||
a.log.InfoContext(parentCtx, "Purged expired signup tokens", slog.Int("count", removed))
|
||||
}
|
||||
|
||||
return a.scheduleCleanup(parentCtx, state)
|
||||
}
|
||||
|
||||
func (a *signupTokenActor) create(parentCtx context.Context, data actor.Envelope) (storedSignupToken, error) {
|
||||
if data == nil {
|
||||
return storedSignupToken{}, fmt.Errorf("request body is empty for method '%s'", signupTokenMethodCreate)
|
||||
}
|
||||
|
||||
var token storedSignupToken
|
||||
err := data.Decode(&token)
|
||||
if err != nil {
|
||||
return storedSignupToken{}, fmt.Errorf("request body is not valid for method '%s': %w", signupTokenMethodCreate, err)
|
||||
}
|
||||
|
||||
state, err := a.mustGetState(parentCtx)
|
||||
if err != nil {
|
||||
return storedSignupToken{}, err
|
||||
}
|
||||
|
||||
state.Tokens[token.Token] = token
|
||||
|
||||
err = a.saveState(parentCtx, state)
|
||||
if err != nil {
|
||||
return storedSignupToken{}, err
|
||||
}
|
||||
|
||||
err = a.scheduleCleanup(parentCtx, state)
|
||||
if err != nil {
|
||||
return storedSignupToken{}, err
|
||||
}
|
||||
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func (a *signupTokenActor) delete(parentCtx context.Context, data actor.Envelope) error {
|
||||
if data == nil {
|
||||
return fmt.Errorf("request body is empty for method '%s'", signupTokenMethodDelete)
|
||||
}
|
||||
|
||||
var req signupTokenDeleteRequest
|
||||
err := data.Decode(&req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("request body is not valid for method '%s': %w", signupTokenMethodDelete, err)
|
||||
}
|
||||
|
||||
state, err := a.mustGetState(parentCtx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Tokens are keyed by their value, so find the one matching the given ID
|
||||
var deleted bool
|
||||
for k, t := range state.Tokens {
|
||||
if t.ID == req.ID {
|
||||
delete(state.Tokens, k)
|
||||
deleted = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !deleted {
|
||||
return nil
|
||||
}
|
||||
|
||||
err = a.saveState(parentCtx, state)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return a.scheduleCleanup(parentCtx, state)
|
||||
}
|
||||
|
||||
func (a *signupTokenActor) consume(parentCtx context.Context, data actor.Envelope) (signupTokenConsumeResponse, error) {
|
||||
if data == nil {
|
||||
return signupTokenConsumeResponse{}, fmt.Errorf("request body is empty for method '%s'", signupTokenMethodConsume)
|
||||
}
|
||||
|
||||
var req signupTokenConsumeRequest
|
||||
err := data.Decode(&req)
|
||||
if err != nil {
|
||||
return signupTokenConsumeResponse{}, fmt.Errorf("request body is not valid for method '%s': %w", signupTokenMethodConsume, err)
|
||||
}
|
||||
|
||||
state, err := a.mustGetState(parentCtx)
|
||||
if err != nil {
|
||||
return signupTokenConsumeResponse{}, err
|
||||
}
|
||||
|
||||
token, ok := state.Tokens[req.Token]
|
||||
switch {
|
||||
case !ok:
|
||||
return signupTokenConsumeResponse{
|
||||
Status: signupTokenConsumeNotFound,
|
||||
}, nil
|
||||
case token.isExpired(time.Now()):
|
||||
return signupTokenConsumeResponse{
|
||||
Status: signupTokenConsumeExpired,
|
||||
}, nil
|
||||
case token.isUsageLimitReached():
|
||||
return signupTokenConsumeResponse{
|
||||
Status: signupTokenConsumeLimitReached,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Atomically consume one use of the token
|
||||
token.UsageCount++
|
||||
state.Tokens[req.Token] = token
|
||||
|
||||
err = a.saveState(parentCtx, state)
|
||||
if err != nil {
|
||||
return signupTokenConsumeResponse{}, err
|
||||
}
|
||||
|
||||
return signupTokenConsumeResponse{
|
||||
Status: signupTokenConsumeOK,
|
||||
UserGroupIDs: token.UserGroupIDs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (a *signupTokenActor) release(parentCtx context.Context, data actor.Envelope) error {
|
||||
if data == nil {
|
||||
return fmt.Errorf("request body is empty for method '%s'", signupTokenMethodRelease)
|
||||
}
|
||||
|
||||
var req signupTokenReleaseRequest
|
||||
err := data.Decode(&req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("request body is not valid for method '%s': %w", signupTokenMethodRelease, err)
|
||||
}
|
||||
|
||||
state, err := a.mustGetState(parentCtx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
token, ok := state.Tokens[req.Token]
|
||||
if !ok || token.UsageCount <= 0 {
|
||||
// The token is gone (for example, expired and purged) or was never consumed: nothing to compensate
|
||||
return nil
|
||||
}
|
||||
|
||||
token.UsageCount--
|
||||
state.Tokens[req.Token] = token
|
||||
|
||||
return a.saveState(parentCtx, state)
|
||||
}
|
||||
|
||||
func (a *signupTokenActor) replace(parentCtx context.Context, data actor.Envelope) error {
|
||||
if data == nil {
|
||||
return fmt.Errorf("request body is empty for method '%s'", signupTokenMethodReplace)
|
||||
}
|
||||
|
||||
var req signupTokenReplaceRequest
|
||||
err := data.Decode(&req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("request body is not valid for method '%s': %w", signupTokenMethodReplace, err)
|
||||
}
|
||||
|
||||
state := &signupTokenActorState{
|
||||
Tokens: make(map[string]storedSignupToken, len(req.Tokens)),
|
||||
}
|
||||
for _, t := range req.Tokens {
|
||||
state.Tokens[t.Token] = t
|
||||
}
|
||||
|
||||
err = a.saveState(parentCtx, state)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return a.scheduleCleanup(parentCtx, state)
|
||||
}
|
||||
|
||||
// scheduleCleanup sets the cleanup alarm to fire when the earliest-expiring token expires, or deletes it when there are no tokens left.
|
||||
func (a *signupTokenActor) scheduleCleanup(parentCtx context.Context, state *signupTokenActorState) (err error) {
|
||||
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
earliest, found := state.earliestExpiration()
|
||||
if !found {
|
||||
// No tokens: remove the alarm (if any)
|
||||
err = a.client.DeleteAlarm(ctx, cleanupAlarmName)
|
||||
if err != nil && !errors.Is(err, actor.ErrAlarmNotFound) {
|
||||
return fmt.Errorf("error deleting cleanup alarm: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
err = a.client.SetAlarm(ctx, cleanupAlarmName, actor.AlarmProperties{DueTime: earliest})
|
||||
if err != nil {
|
||||
return fmt.Errorf("error setting cleanup alarm: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// mustGetState retrieves the actor state, initializing an empty one if it doesn't exist yet.
|
||||
func (a *signupTokenActor) mustGetState(parentCtx context.Context) (*signupTokenActorState, error) {
|
||||
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
state, err := a.client.GetState(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error retrieving actor state: %w", err)
|
||||
}
|
||||
|
||||
if state == nil {
|
||||
state = &signupTokenActorState{
|
||||
Tokens: map[string]storedSignupToken{},
|
||||
}
|
||||
} else if state.Tokens == nil {
|
||||
state.Tokens = map[string]storedSignupToken{}
|
||||
}
|
||||
|
||||
return state, nil
|
||||
}
|
||||
|
||||
func (a *signupTokenActor) saveState(parentCtx context.Context, state *signupTokenActorState) error {
|
||||
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
|
||||
defer cancel()
|
||||
err := a.client.SetState(ctx, state, nil)
|
||||
if err != nil {
|
||||
return fmt.Errorf("error saving actor state: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// collectTokens returns all tokens in the state as a slice.
|
||||
func collectTokens(state *signupTokenActorState) []storedSignupToken {
|
||||
if state == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
tokens := make([]storedSignupToken, len(state.Tokens))
|
||||
var i int
|
||||
for _, t := range state.Tokens {
|
||||
tokens[i] = t
|
||||
i++
|
||||
}
|
||||
|
||||
return tokens
|
||||
}
|
||||
194
backend/internal/usersignup/signuptoken_actor_test.go
Normal file
194
backend/internal/usersignup/signuptoken_actor_test.go
Normal file
@@ -0,0 +1,194 @@
|
||||
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 singleton actor registered and returns its service.
|
||||
// seed, if not nil, is used as the bootstrap data to migrate tokens into the actor's state.
|
||||
func newSignupTokenActorService(t *testing.T, seed []storedSignupToken) *actor.Service {
|
||||
t.Helper()
|
||||
|
||||
var svc *actor.Service
|
||||
testutils.NewActorHostForTest(t, func(t *testing.T, h *local.Host) {
|
||||
err := h.RegisterSingletonActor(
|
||||
SignupTokenActorType, NewSignupTokenActor,
|
||||
local.WithBootstrapData(&signupTokenBootstrap{Tokens: seed}),
|
||||
local.WithIdleTimeout(-1),
|
||||
)
|
||||
require.NoError(t, err)
|
||||
svc = h.Service()
|
||||
})
|
||||
require.NotNil(t, svc)
|
||||
return svc
|
||||
}
|
||||
|
||||
func createSignupTokenForTest(t *testing.T, svc *actor.Service, token storedSignupToken) {
|
||||
t.Helper()
|
||||
_, err := svc.Invoke(t.Context(), SignupTokenActorType, actor.SingletonActorID, signupTokenMethodCreate, token)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func consumeSignupTokenForTest(t *testing.T, svc *actor.Service, token string) signupTokenConsumeResponse {
|
||||
t.Helper()
|
||||
res, err := svc.Invoke(t.Context(), SignupTokenActorType, actor.SingletonActorID, signupTokenMethodConsume, signupTokenConsumeRequest{Token: token})
|
||||
require.NoError(t, err)
|
||||
var out signupTokenConsumeResponse
|
||||
require.NoError(t, res.Decode(&out))
|
||||
return out
|
||||
}
|
||||
|
||||
func listSignupTokensForTest(t *testing.T, svc *actor.Service) []storedSignupToken {
|
||||
t.Helper()
|
||||
res, err := svc.Peek(t.Context(), SignupTokenActorType, actor.SingletonActorID, signupTokenMethodList, nil)
|
||||
require.NoError(t, err)
|
||||
var out signupTokenListResponse
|
||||
require.NoError(t, res.Decode(&out))
|
||||
return out.Tokens
|
||||
}
|
||||
|
||||
func TestSignupTokenActorConsume(t *testing.T) {
|
||||
svc := newSignupTokenActorService(t, nil)
|
||||
|
||||
createSignupTokenForTest(t, svc, storedSignupToken{
|
||||
ID: "id-1",
|
||||
Token: "token-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, nil)
|
||||
|
||||
res := consumeSignupTokenForTest(t, svc, "does-not-exist")
|
||||
require.Equal(t, signupTokenConsumeNotFound, res.Status)
|
||||
}
|
||||
|
||||
func TestSignupTokenActorConsumeExpired(t *testing.T) {
|
||||
svc := newSignupTokenActorService(t, nil)
|
||||
|
||||
createSignupTokenForTest(t, svc, storedSignupToken{
|
||||
ID: "id-expired",
|
||||
Token: "token-expired",
|
||||
ExpiresAt: time.Now().Add(-time.Minute),
|
||||
UsageLimit: 1,
|
||||
CreatedAt: time.Now().Add(-time.Hour),
|
||||
})
|
||||
|
||||
res := consumeSignupTokenForTest(t, svc, "token-expired")
|
||||
require.Equal(t, signupTokenConsumeExpired, res.Status)
|
||||
}
|
||||
|
||||
func TestSignupTokenActorRelease(t *testing.T) {
|
||||
svc := newSignupTokenActorService(t, nil)
|
||||
|
||||
createSignupTokenForTest(t, svc, storedSignupToken{
|
||||
ID: "id-2",
|
||||
Token: "token-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, actor.SingletonActorID, signupTokenMethodRelease, signupTokenReleaseRequest{Token: "token-2"})
|
||||
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, nil)
|
||||
|
||||
createSignupTokenForTest(t, svc, storedSignupToken{
|
||||
ID: "id-3",
|
||||
Token: "token-3",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
UsageLimit: 1,
|
||||
CreatedAt: time.Now(),
|
||||
})
|
||||
require.Len(t, listSignupTokensForTest(t, svc), 1)
|
||||
|
||||
_, err := svc.Invoke(t.Context(), SignupTokenActorType, actor.SingletonActorID, signupTokenMethodDelete, signupTokenDeleteRequest{ID: "id-3"})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Empty(t, listSignupTokensForTest(t, svc))
|
||||
|
||||
// The token can no longer be consumed
|
||||
require.Equal(t, signupTokenConsumeNotFound, consumeSignupTokenForTest(t, svc, "token-3").Status)
|
||||
}
|
||||
|
||||
func TestSignupTokenActorBootstrapMigration(t *testing.T) {
|
||||
now := time.Now()
|
||||
seed := []storedSignupToken{
|
||||
{ID: "valid", Token: "valid-token", ExpiresAt: now.Add(time.Hour), UsageLimit: 1, CreatedAt: now},
|
||||
{ID: "expired", Token: "expired-token", ExpiresAt: now.Add(-time.Hour), UsageLimit: 1, CreatedAt: now},
|
||||
}
|
||||
svc := newSignupTokenActorService(t, seed)
|
||||
|
||||
// The singleton actor bootstraps asynchronously once the host is ready. Wait until the migrated (non-expired) token is available.
|
||||
require.Eventually(t, func() bool {
|
||||
tokens := listSignupTokensForTest(t, svc)
|
||||
return len(tokens) == 1 && tokens[0].Token == "valid-token"
|
||||
}, 10*time.Second, 20*time.Millisecond, "signup token actor was not bootstrapped in time")
|
||||
|
||||
// The migrated token can be consumed
|
||||
require.Equal(t, signupTokenConsumeOK, consumeSignupTokenForTest(t, svc, "valid-token").Status)
|
||||
}
|
||||
|
||||
func TestSignupTokenStateRemoveExpired(t *testing.T) {
|
||||
now := time.Now()
|
||||
state := &signupTokenActorState{Tokens: map[string]storedSignupToken{
|
||||
"a": {Token: "a", ExpiresAt: now.Add(time.Hour)},
|
||||
"b": {Token: "b", ExpiresAt: now.Add(-time.Minute)},
|
||||
"c": {Token: "c", ExpiresAt: now.Add(-time.Hour)},
|
||||
}}
|
||||
|
||||
removed := state.removeExpired(now)
|
||||
require.Equal(t, 2, removed)
|
||||
require.Len(t, state.Tokens, 1)
|
||||
_, ok := state.Tokens["a"]
|
||||
require.True(t, ok)
|
||||
}
|
||||
|
||||
func TestSignupTokenStateEarliestExpiration(t *testing.T) {
|
||||
now := time.Now()
|
||||
|
||||
empty := &signupTokenActorState{Tokens: map[string]storedSignupToken{}}
|
||||
_, found := empty.earliestExpiration()
|
||||
require.False(t, found)
|
||||
|
||||
state := &signupTokenActorState{Tokens: map[string]storedSignupToken{
|
||||
"a": {Token: "a", ExpiresAt: now.Add(2 * time.Hour)},
|
||||
"b": {Token: "b", ExpiresAt: now.Add(time.Hour)},
|
||||
"c": {Token: "c", ExpiresAt: now.Add(3 * time.Hour)},
|
||||
}}
|
||||
earliest, found := state.earliestExpiration()
|
||||
require.True(t, found)
|
||||
require.Equal(t, now.Add(time.Hour), earliest)
|
||||
}
|
||||
48
backend/internal/usersignup/testing_e2etest.go
Normal file
48
backend/internal/usersignup/testing_e2etest.go
Normal file
@@ -0,0 +1,48 @@
|
||||
//go:build e2etest
|
||||
|
||||
package usersignup
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/italypaleale/francis/actor"
|
||||
"github.com/italypaleale/francis/host/local"
|
||||
)
|
||||
|
||||
// SignupTokenSeed describes a signup token to seed into the actor state, used by E2E test setup.
|
||||
type SignupTokenSeed struct {
|
||||
ID string
|
||||
Token string
|
||||
ExpiresAt time.Time
|
||||
UsageLimit int
|
||||
UsageCount int
|
||||
UserGroupIDs []string
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
// SeedSignupTokens replaces the signup token singleton actor's state with the given tokens.
|
||||
// It's intended for E2E test setup, where fixture tokens need to be created with exact values.
|
||||
func SeedSignupTokens(ctx context.Context, actors *local.Host, seeds []SignupTokenSeed) error {
|
||||
tokens := make([]storedSignupToken, len(seeds))
|
||||
for i, s := range seeds {
|
||||
tokens[i] = storedSignupToken{
|
||||
ID: s.ID,
|
||||
Token: s.Token,
|
||||
ExpiresAt: s.ExpiresAt,
|
||||
UsageLimit: s.UsageLimit,
|
||||
UsageCount: s.UsageCount,
|
||||
UserGroupIDs: s.UserGroupIDs,
|
||||
CreatedAt: s.CreatedAt,
|
||||
}
|
||||
}
|
||||
|
||||
_, err := actors.Service().Invoke(ctx, SignupTokenActorType, actor.SingletonActorID, signupTokenMethodReplace, signupTokenReplaceRequest{
|
||||
Tokens: tokens,
|
||||
})
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to seed signup tokens into actor: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
-- Recreate the one_time_access_tokens table with the schema it had before it was dropped.
|
||||
CREATE TABLE one_time_access_tokens
|
||||
(
|
||||
id UUID NOT NULL PRIMARY KEY,
|
||||
created_at TIMESTAMPTZ,
|
||||
token VARCHAR(255) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMPTZ NOT NULL,
|
||||
user_id UUID NOT NULL REFERENCES users ON DELETE CASCADE,
|
||||
device_token VARCHAR(16)
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_one_time_access_tokens_expires_at ON one_time_access_tokens (expires_at);
|
||||
|
||||
-- Recreate the signup token tables with the schema they had before they were frozen.
|
||||
CREATE TABLE signup_tokens (
|
||||
id UUID NOT NULL PRIMARY KEY,
|
||||
created_at TIMESTAMPTZ NOT NULL,
|
||||
token VARCHAR(255) NOT NULL UNIQUE,
|
||||
expires_at TIMESTAMPTZ NOT NULL,
|
||||
usage_limit INTEGER NOT NULL DEFAULT 1,
|
||||
usage_count INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE INDEX idx_signup_tokens_token ON signup_tokens(token);
|
||||
CREATE INDEX idx_signup_tokens_expires_at ON signup_tokens(expires_at);
|
||||
|
||||
CREATE TABLE signup_tokens_user_groups
|
||||
(
|
||||
signup_token_id UUID NOT NULL,
|
||||
user_group_id UUID NOT NULL,
|
||||
PRIMARY KEY (signup_token_id, user_group_id),
|
||||
FOREIGN KEY (signup_token_id) REFERENCES signup_tokens (id) ON DELETE CASCADE,
|
||||
FOREIGN KEY (user_group_id) REFERENCES user_groups (id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
-- Restore the signup tokens from the frozen JSON document stored in the "kv" table.
|
||||
-- json_array_elements expands the JSON array into one row per token object.
|
||||
INSERT INTO signup_tokens (id, created_at, token, expires_at, usage_limit, usage_count)
|
||||
SELECT
|
||||
(e ->> 'id')::uuid,
|
||||
to_timestamp((e ->> 'createdAt')::bigint),
|
||||
e ->> 'token',
|
||||
to_timestamp((e ->> 'expiresAt')::bigint),
|
||||
(e ->> 'usageLimit')::int,
|
||||
(e ->> 'usageCount')::int
|
||||
FROM kv, json_array_elements(kv."value"::json) AS e
|
||||
WHERE kv."key" = 'signup_tokens_migrated';
|
||||
|
||||
-- Restore the token/user-group associations, expanding each token's nested userGroupIds array.
|
||||
INSERT INTO signup_tokens_user_groups (signup_token_id, user_group_id)
|
||||
SELECT
|
||||
(e ->> 'id')::uuid,
|
||||
g.value::uuid
|
||||
FROM kv, json_array_elements(kv."value"::json) AS e, json_array_elements_text(e -> 'userGroupIds') AS g
|
||||
WHERE kv."key" = 'signup_tokens_migrated';
|
||||
|
||||
-- Remove the frozen signup tokens from the "kv" table.
|
||||
DELETE FROM kv WHERE "key" = 'signup_tokens_migrated';
|
||||
@@ -0,0 +1,28 @@
|
||||
-- One-time access tokens are now stored in the actor state store, so the table is no longer needed.
|
||||
DROP TABLE IF EXISTS one_time_access_tokens;
|
||||
|
||||
-- Freeze the signup tokens.
|
||||
-- Encode every signup token (with its user group IDs) as a single JSON array and store it in the "kv" table under the "signup_tokens_migrated" key, so the singleton signup token actor can seed its state from it on first startup.
|
||||
-- The "HAVING count(*) > 0" clause ensures nothing is written to the "kv" table when there are no signup tokens.
|
||||
-- Timestamps are stored as Unix seconds so the frozen format is identical across databases.
|
||||
INSERT INTO kv ("key", "value")
|
||||
SELECT 'signup_tokens_migrated', json_agg(
|
||||
json_build_object(
|
||||
'id', st.id,
|
||||
'token', st.token,
|
||||
'expiresAt', extract(epoch FROM st.expires_at)::bigint,
|
||||
'usageLimit', st.usage_limit,
|
||||
'usageCount', st.usage_count,
|
||||
'createdAt', extract(epoch FROM st.created_at)::bigint,
|
||||
'userGroupIds', COALESCE(
|
||||
(SELECT json_agg(stug.user_group_id) FROM signup_tokens_user_groups stug WHERE stug.signup_token_id = st.id),
|
||||
'[]'::json
|
||||
)
|
||||
)
|
||||
)::text
|
||||
FROM signup_tokens st
|
||||
HAVING count(*) > 0;
|
||||
|
||||
-- Drop the now-frozen signup token tables.
|
||||
DROP TABLE signup_tokens_user_groups;
|
||||
DROP TABLE signup_tokens;
|
||||
@@ -0,0 +1,62 @@
|
||||
PRAGMA foreign_keys=OFF;
|
||||
BEGIN;
|
||||
|
||||
-- Recreate the one_time_access_tokens table with the schema it had before it was dropped.
|
||||
CREATE TABLE one_time_access_tokens
|
||||
(
|
||||
id TEXT PRIMARY KEY,
|
||||
created_at DATETIME NOT NULL,
|
||||
token TEXT NOT NULL UNIQUE,
|
||||
expires_at DATETIME NOT NULL,
|
||||
user_id TEXT NOT NULL REFERENCES users ON DELETE CASCADE,
|
||||
device_token TEXT
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS idx_one_time_access_tokens_expires_at ON one_time_access_tokens (expires_at);
|
||||
|
||||
-- Recreate the signup token tables with the schema they had before they were frozen.
|
||||
CREATE TABLE signup_tokens (
|
||||
id TEXT NOT NULL PRIMARY KEY,
|
||||
created_at DATETIME NOT NULL,
|
||||
token TEXT NOT NULL UNIQUE,
|
||||
expires_at DATETIME NOT NULL,
|
||||
usage_limit INTEGER NOT NULL DEFAULT 1,
|
||||
usage_count INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE INDEX idx_signup_tokens_token ON signup_tokens(token);
|
||||
CREATE INDEX idx_signup_tokens_expires_at ON signup_tokens(expires_at);
|
||||
|
||||
CREATE TABLE signup_tokens_user_groups
|
||||
(
|
||||
signup_token_id TEXT NOT NULL,
|
||||
user_group_id TEXT NOT NULL,
|
||||
PRIMARY KEY (signup_token_id, user_group_id),
|
||||
FOREIGN KEY (signup_token_id) REFERENCES signup_tokens (id) ON DELETE CASCADE,
|
||||
FOREIGN KEY (user_group_id) REFERENCES user_groups (id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
-- Restore the signup tokens from the frozen JSON document stored in the "kv" table.
|
||||
-- json_each expands the JSON array into one row per token object.
|
||||
INSERT INTO signup_tokens (id, created_at, token, expires_at, usage_limit, usage_count)
|
||||
SELECT
|
||||
json_extract(e.value, '$.id'),
|
||||
json_extract(e.value, '$.createdAt'),
|
||||
json_extract(e.value, '$.token'),
|
||||
json_extract(e.value, '$.expiresAt'),
|
||||
json_extract(e.value, '$.usageLimit'),
|
||||
json_extract(e.value, '$.usageCount')
|
||||
FROM kv, json_each(kv."value") AS e
|
||||
WHERE kv."key" = 'signup_tokens_migrated';
|
||||
|
||||
-- Restore the token/user-group associations, expanding each token's nested userGroupIds array.
|
||||
INSERT INTO signup_tokens_user_groups (signup_token_id, user_group_id)
|
||||
SELECT
|
||||
json_extract(e.value, '$.id'),
|
||||
g.value
|
||||
FROM kv, json_each(kv."value") AS e, json_each(json_extract(e.value, '$.userGroupIds')) AS g
|
||||
WHERE kv."key" = 'signup_tokens_migrated';
|
||||
|
||||
-- Remove the frozen signup tokens from the "kv" table.
|
||||
DELETE FROM kv WHERE "key" = 'signup_tokens_migrated';
|
||||
|
||||
COMMIT;
|
||||
PRAGMA foreign_keys=ON;
|
||||
@@ -0,0 +1,35 @@
|
||||
PRAGMA foreign_keys=OFF;
|
||||
BEGIN;
|
||||
|
||||
-- One-time access tokens are now stored in the actor state store, so the table is no longer needed.
|
||||
DROP TABLE IF EXISTS one_time_access_tokens;
|
||||
|
||||
-- Freeze the signup tokens.
|
||||
-- Encode every signup token (with its user group IDs) as a single JSON array and store it in the "kv" table under the "signup_tokens_migrated" key, so the singleton signup token actor can seed its state from it on first startup.
|
||||
-- The "HAVING count(*) > 0" clause ensures nothing is written to the "kv" table when there are no signup tokens.
|
||||
-- Timestamps are stored as Unix seconds, matching how DateTime values are persisted on SQLite.
|
||||
INSERT INTO kv ("key", "value")
|
||||
SELECT 'signup_tokens_migrated', json_group_array(
|
||||
json_object(
|
||||
'id', st.id,
|
||||
'token', st.token,
|
||||
'expiresAt', st.expires_at,
|
||||
'usageLimit', st.usage_limit,
|
||||
'usageCount', st.usage_count,
|
||||
'createdAt', st.created_at,
|
||||
'userGroupIds', json((
|
||||
SELECT COALESCE(json_group_array(stug.user_group_id), json_array())
|
||||
FROM signup_tokens_user_groups stug
|
||||
WHERE stug.signup_token_id = st.id
|
||||
))
|
||||
)
|
||||
)
|
||||
FROM signup_tokens st
|
||||
HAVING count(*) > 0;
|
||||
|
||||
-- Drop the now-frozen signup token tables.
|
||||
DROP TABLE signup_tokens_user_groups;
|
||||
DROP TABLE signup_tokens;
|
||||
|
||||
COMMIT;
|
||||
PRAGMA foreign_keys=ON;
|
||||
5429
frontend/package-lock.json
generated
Normal file
5429
frontend/package-lock.json
generated
Normal file
File diff suppressed because it is too large
Load Diff
@@ -238,72 +238,6 @@
|
||||
"user_group_id": "c7ae7c01-28a3-4f3c-9572-1ee734ea8368"
|
||||
}
|
||||
],
|
||||
"one_time_access_tokens": [
|
||||
{
|
||||
"created_at": "2025-11-25T12:39:02Z",
|
||||
"expires_at": "2025-11-25T13:39:02Z",
|
||||
"id": "bf877753-4ea4-4c9c-bbbd-e198bb201cb8",
|
||||
"token": "HPe6k6uiDRRVuAQV",
|
||||
"device_token": null,
|
||||
"user_id": "f4b89dc2-62fb-46bf-9f5f-c34f4eafe93e"
|
||||
},
|
||||
{
|
||||
"created_at": "2025-11-25T12:39:02Z",
|
||||
"expires_at": "2025-11-25T12:39:01Z",
|
||||
"id": "d3afae24-fe2d-4a98-abec-cf0b8525096a",
|
||||
"token": "YCGDtftvsvYWiXd0",
|
||||
"device_token": null,
|
||||
"user_id": "f4b89dc2-62fb-46bf-9f5f-c34f4eafe93e"
|
||||
},
|
||||
{
|
||||
"created_at": "2025-11-25T12:39:02Z",
|
||||
"expires_at": "2025-11-25T13:39:02Z",
|
||||
"id": "defd5164-9d9b-4228-bbce-708e33f49360",
|
||||
"token": "one-time-token",
|
||||
"device_token": null,
|
||||
"user_id": "f4b89dc2-62fb-46bf-9f5f-c34f4eafe93e"
|
||||
}
|
||||
],
|
||||
"signup_tokens": [
|
||||
{
|
||||
"created_at": "2025-11-25T12:39:02Z",
|
||||
"expires_at": "2025-11-26T12:39:02Z",
|
||||
"id": "a1b2c3d4-e5f6-7890-abcd-ef1234567890",
|
||||
"token": "VALID1234567890A",
|
||||
"usage_count": 0,
|
||||
"usage_limit": 1
|
||||
},
|
||||
{
|
||||
"created_at": "2025-11-25T12:39:02Z",
|
||||
"expires_at": "2025-12-02T12:39:02Z",
|
||||
"id": "dc3c9c96-714e-48eb-926e-2d7c7858e6cf",
|
||||
"token": "PARTIAL567890ABC",
|
||||
"usage_count": 2,
|
||||
"usage_limit": 5
|
||||
},
|
||||
{
|
||||
"created_at": "2025-11-25T12:39:02Z",
|
||||
"expires_at": "2025-11-24T12:39:02Z",
|
||||
"id": "44de1863-ffa5-4db1-9507-4887cd7a1e3f",
|
||||
"token": "EXPIRED34567890B",
|
||||
"usage_count": 1,
|
||||
"usage_limit": 3
|
||||
},
|
||||
{
|
||||
"created_at": "2025-11-25T12:39:02Z",
|
||||
"expires_at": "2025-11-26T12:39:02Z",
|
||||
"id": "f1b1678b-7720-4d8b-8f91-1dbff1e2d02b",
|
||||
"token": "FULLYUSED567890C",
|
||||
"usage_count": 1,
|
||||
"usage_limit": 1
|
||||
}
|
||||
],
|
||||
"signup_tokens_user_groups": [
|
||||
{
|
||||
"signup_token_id": "a1b2c3d4-e5f6-7890-abcd-ef1234567890",
|
||||
"user_group_id": "c7ae7c01-28a3-4f3c-9572-1ee734ea8368"
|
||||
}
|
||||
],
|
||||
"user_authorized_oidc_clients": [
|
||||
{
|
||||
"client_id": "3654a746-35d4-4321-ac61-0bdcff2b4055",
|
||||
|
||||
Reference in New Issue
Block a user