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

Co-authored-by: Elias Schneider <login@eliasschneider.com>
This commit is contained in:
Alessandro (Ale) Segala
2026-07-26 15:32:43 +02:00
committed by GitHub
co-authored by Elias Schneider
parent 531bb5f0cf
commit a1b4e1d2b2
37 changed files with 7989 additions and 905 deletions
+167
View File
@@ -0,0 +1,167 @@
package onetimeaccess
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).
// TokenActorType is the actor type for the one-time access token actor
const TokenActorType = "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 (
// TokenMethodRestore stores a token's state, and is also how a consumed token is put back
TokenMethodRestore = "restore"
tokenMethodConsume = "consume"
)
// tokenConsumeStatus is the outcome of a "consume" invocation.
type tokenConsumeStatus string
const (
// tokenConsumeOK indicates the token was valid and has been consumed
tokenConsumeOK tokenConsumeStatus = "ok"
// tokenConsumeNotFound indicates the token doesn't exist (or has expired)
tokenConsumeNotFound tokenConsumeStatus = "not_found"
// tokenConsumeDeviceMismatch indicates the provided device token doesn't match
tokenConsumeDeviceMismatch tokenConsumeStatus = "device_mismatch"
)
// TokenState is the persisted state of a one-time access token actor.
// The token value itself is the actor's ID, so it isn't repeated here.
type TokenState struct {
UserID string
DeviceToken *string
ExpiresAt time.Time
}
// tokenConsumeRequest is the payload for the "consume" method
type tokenConsumeRequest struct {
DeviceToken string
}
// tokenConsumeResponse is the response of the "consume" method
type tokenConsumeResponse struct {
Status tokenConsumeStatus
// State is included only when Status is "ok", so the caller can restore it if a later step fails
State TokenState
}
// tokenActor is the actor that manages a single one-time access token
type tokenActor struct {
log *slog.Logger
client actor.Client[TokenState]
}
// NewTokenActor allocates a new one-time access token actor
// It satisfies actor.Factory
func NewTokenActor(actorID string, service *actor.Service) actor.Actor {
return &tokenActor{
log: slog.With(
slog.String("scope", "actor"),
slog.String("actorType", TokenActorType),
),
client: actor.NewActorClient[TokenState](TokenActorType, actorID, service),
}
}
// Invoke implements actor.ActorInvoke
func (a *tokenActor) Invoke(parentCtx context.Context, method string, data actor.Envelope) (any, error) {
switch method {
case tokenMethodConsume:
return a.consume(parentCtx, data)
case TokenMethodRestore:
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 *tokenActor) consume(parentCtx context.Context, data actor.Envelope) (tokenConsumeResponse, error) {
var req tokenConsumeRequest
if data != nil {
err := data.Decode(&req)
if err != nil {
return tokenConsumeResponse{}, fmt.Errorf("request body is not valid for method '%s': %w", tokenMethodConsume, err)
}
}
ctx, cancel := context.WithTimeout(parentCtx, 10*time.Second)
defer cancel()
state, err := a.client.GetState(ctx)
if err != nil {
return tokenConsumeResponse{}, 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 tokenConsumeResponse{
Status: tokenConsumeNotFound,
}, 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 tokenConsumeResponse{
Status: tokenConsumeDeviceMismatch,
}, 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 tokenConsumeResponse{}, fmt.Errorf("error deleting actor state: %w", err)
}
return tokenConsumeResponse{
Status: tokenConsumeOK,
State: state,
}, nil
}
// restore re-creates the token state, used to compensate when a step after consuming the token fails.
func (a *tokenActor) restore(parentCtx context.Context, data actor.Envelope) error {
if data == nil {
return fmt.Errorf("request body is empty for method '%s'", TokenMethodRestore)
}
var state TokenState
err := data.Decode(&state)
if err != nil {
return fmt.Errorf("request body is not valid for method '%s': %w", TokenMethodRestore, 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
}
+18
View File
@@ -0,0 +1,18 @@
package onetimeaccess
import (
"github.com/pocket-id/pocket-id/backend/internal/utils"
)
type tokenCreateDto struct {
TTL utils.JSONDuration `json:"ttl" binding:"ttl"`
}
type emailAsUnauthenticatedUserDto struct {
Email string `json:"email" binding:"required,email" unorm:"nfc"`
RedirectPath string `json:"redirectPath"`
}
type emailAsAdminDto struct {
TTL utils.JSONDuration `json:"ttl" binding:"ttl"`
}
+198
View File
@@ -0,0 +1,198 @@
package onetimeaccess
import (
"fmt"
"net/http"
"time"
"github.com/gin-gonic/gin"
"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/utils/cookie"
)
const defaultTokenDuration = 15 * time.Minute
type handler struct {
service *Service
appConfig AppConfigResolver
}
func newHandler(service *Service, appConfig AppConfigResolver) *handler {
return &handler{service: service, appConfig: appConfig}
}
func (h *handler) createToken(c *gin.Context, own bool) {
var input tokenCreateDto
err := c.ShouldBindJSON(&input)
if err != nil {
_ = c.Error(err)
return
}
var (
userID string
ttl time.Duration
)
if own {
// Get user ID from context and force the default TTL
userID = c.GetString("userID")
ttl = defaultTokenDuration
} else {
// Get user ID from URL parameter, and optional TTL from body
userID = c.Param("id")
ttl = input.TTL.Duration
if ttl <= 0 {
ttl = defaultTokenDuration
}
}
if userID == "" {
_ = c.Error(&common.UserIdNotProvidedError{})
return
}
token, err := h.service.CreateToken(c.Request.Context(), userID, ttl)
if err != nil {
_ = c.Error(err)
return
}
c.JSON(http.StatusCreated, gin.H{"token": token})
}
// createOwnToken godoc
// @Summary Create one-time access token for current user
// @Description Generate a one-time access token for the currently authenticated user
// @Tags Users
// @Param body body tokenCreateDto true "Token options"
// @Success 201 {object} object "{ \"token\": \"string\" }"
// @Router /api/users/me/one-time-access-token [post]
func (h *handler) createOwnToken(c *gin.Context) {
h.createToken(c, true)
}
// createTokenForUser godoc
// @Summary Create one-time access token for user (admin)
// @Description Generate a one-time access token for a specific user (admin only)
// @Tags Users
// @Param id path string true "User ID"
// @Param body body tokenCreateDto true "Token options"
// @Success 201 {object} object "{ \"token\": \"string\" }"
// @Router /api/users/{id}/one-time-access-token [post]
func (h *handler) createTokenForUser(c *gin.Context) {
h.createToken(c, false)
}
// requestEmailAsUnauthenticatedUser godoc
// @Summary Request one-time access email
// @Description Request a one-time access email for unauthenticated users
// @Tags Users
// @Accept json
// @Produce json
// @Param body body emailAsUnauthenticatedUserDto true "Email request information"
// @Success 204 "No Content"
// @Router /api/one-time-access-email [post]
func (h *handler) requestEmailAsUnauthenticatedUser(c *gin.Context) {
dbConfig, err := h.appConfig.GetConfig(c.Request.Context())
if err != nil {
_ = c.Error(fmt.Errorf("error loading app configuration: %w", err))
return
}
var input emailAsUnauthenticatedUserDto
err = dto.ShouldBindWithNormalizedJSON(c, &input)
if err != nil {
_ = c.Error(err)
return
}
deviceToken, err := h.service.RequestOneTimeAccessEmailAsUnauthenticatedUser(c.Request.Context(), dbConfig, input.Email, input.RedirectPath)
if err != nil {
_ = c.Error(err)
return
}
cookie.AddDeviceTokenCookie(c, deviceToken)
c.Status(http.StatusNoContent)
}
// requestEmailAsAdmin godoc
// @Summary Request one-time access email (admin)
// @Description Request a one-time access email for a specific user (admin only)
// @Tags Users
// @Accept json
// @Produce json
// @Param id path string true "User ID"
// @Param body body emailAsAdminDto true "Email request options"
// @Success 204 "No Content"
// @Router /api/users/{id}/one-time-access-email [post]
func (h *handler) requestEmailAsAdmin(c *gin.Context) {
dbConfig, err := h.appConfig.GetConfig(c.Request.Context())
if err != nil {
_ = c.Error(fmt.Errorf("error loading app configuration: %w", err))
return
}
var input emailAsAdminDto
err = c.ShouldBindJSON(&input)
if err != nil {
_ = c.Error(err)
return
}
userID := c.Param("id")
ttl := input.TTL.Duration
if ttl <= 0 {
ttl = defaultTokenDuration
}
err = h.service.RequestOneTimeAccessEmailAsAdmin(c.Request.Context(), dbConfig, userID, ttl)
if err != nil {
_ = c.Error(err)
return
}
c.Status(http.StatusNoContent)
}
// exchangeToken godoc
// @Summary Exchange one-time access token
// @Description Exchange a one-time access token for a session token
// @Tags Users
// @Param token path string true "One-time access token"
// @Success 200 {object} dto.UserDto
// @Router /api/one-time-access-token/{token} [post]
func (h *handler) exchangeToken(c *gin.Context) {
cfg, err := h.appConfig.GetConfig(c.Request.Context())
if err != nil {
_ = c.Error(fmt.Errorf("error loading app configuration: %w", err))
return
}
loginCode := c.Param("token")
// reject invalid length login codes
if len(loginCode) != 6 && len(loginCode) != 16 {
_ = c.Error(&common.TokenInvalidOrExpiredError{})
return
}
deviceToken, _ := c.Cookie(cookie.DeviceTokenCookieName)
user, token, err := h.service.ExchangeToken(c.Request.Context(), cfg, loginCode, deviceToken, c.ClientIP(), c.Request.UserAgent())
if err != nil {
_ = c.Error(err)
return
}
var userDto dto.UserDto
err = dto.MapStruct(user, &userDto)
if err != nil {
_ = c.Error(err)
return
}
maxAge := int(cfg.SessionDuration.AsDurationMinutes().Seconds())
cookie.AddAccessTokenCookie(c, maxAge, token)
c.JSON(http.StatusOK, userDto)
}
+86
View File
@@ -0,0 +1,86 @@
package onetimeaccess
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"
"github.com/pocket-id/pocket-id/backend/internal/model"
"github.com/pocket-id/pocket-id/backend/internal/utils/email"
)
// EmailData is the data rendered in the one-time access email
type EmailData struct {
Code string
LoginLink string
LoginLinkWithCode string
ExpirationString string
}
// EmailSender sends the one-time access email
type EmailSender interface {
SendOneTimeAccessEmail(ctx context.Context, dbConfig *appconfig.AppConfigModel, to email.Address, data EmailData) error
}
type TokenService interface {
GenerateAccessToken(user model.User, authenticationMethod string, sessionDuration time.Duration) (string, error)
}
type AuditLogger interface {
Create(ctx context.Context, event model.AuditLogEvent, ipAddress, userAgent, userID string, data model.AuditLogData, tx *gorm.DB) (model.AuditLog, bool)
}
type UserProvider interface {
GetUser(ctx context.Context, userID string) (model.User, error)
}
// AppConfigResolver loads the current application configuration, so handlers can pass it explicitly to the service methods that need it
type AppConfigResolver interface {
GetConfig(ctx context.Context) (*appconfig.AppConfigModel, error)
}
type Dependencies struct {
DB *gorm.DB
Actors *local.Host
Signer TokenService
AuditLog AuditLogger
UserProvider UserProvider
EmailSender EmailSender
AppConfig AppConfigResolver
}
type Module struct {
service *Service
handler *handler
}
func New(deps Dependencies) (*Module, error) {
// Register the actor that manages a one-time access token
// Each token is its own actor, whose actor ID is the token's value
err := deps.Actors.RegisterActor(TokenActorType, NewTokenActor)
if err != nil {
return nil, fmt.Errorf("error registering the %s actor: %w", TokenActorType, err)
}
service := newService(deps, deps.Actors.Service())
return &Module{
service: service,
handler: newHandler(service, deps.AppConfig),
}, nil
}
// RegisterRoutes mounts the one-time access token endpoints
// auth guards the admin routes and ownAuth the current user's own token, while the rate limiters throttle the public exchange and email endpoints
func (m *Module) RegisterRoutes(apiGroup *gin.RouterGroup, auth, ownAuth, exchangeRateLimit, emailRateLimit gin.HandlerFunc) {
apiGroup.POST("/users/me/one-time-access-token", ownAuth, m.handler.createOwnToken)
apiGroup.POST("/users/:id/one-time-access-token", auth, m.handler.createTokenForUser)
apiGroup.POST("/users/:id/one-time-access-email", auth, m.handler.requestEmailAsAdmin)
apiGroup.POST("/one-time-access-token/:token", exchangeRateLimit, m.handler.exchangeToken)
apiGroup.POST("/one-time-access-email", emailRateLimit, m.handler.requestEmailAsUnauthenticatedUser)
}
+283
View File
@@ -0,0 +1,283 @@
package onetimeaccess
import (
"context"
"errors"
"fmt"
"log/slog"
"net/url"
"strings"
"time"
"github.com/italypaleale/francis/actor"
"gorm.io/gorm"
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
"github.com/pocket-id/pocket-id/backend/internal/common"
"github.com/pocket-id/pocket-id/backend/internal/model"
"github.com/pocket-id/pocket-id/backend/internal/utils"
"github.com/pocket-id/pocket-id/backend/internal/utils/email"
)
// authenticationMethodOneTimePassword identifies one-time password/code authentication
// It must match the value emitted by the JWT service in the access token's "amr" claim
const authenticationMethodOneTimePassword = "otp"
// TokenStore 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 TokenStore interface {
SetState(ctx context.Context, actorType string, actorID string, state any, opts *actor.SetStateOpts) error
}
type Service struct {
db *gorm.DB
actorService *actor.Service
userProvider UserProvider
signer TokenService
auditLog AuditLogger
emailSender EmailSender
}
func newService(deps Dependencies, actorService *actor.Service) *Service {
return &Service{
db: deps.DB,
actorService: actorService,
userProvider: deps.UserProvider,
signer: deps.Signer,
auditLog: deps.AuditLog,
emailSender: deps.EmailSender,
}
}
func (s *Service) RequestOneTimeAccessEmailAsAdmin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID string, ttl time.Duration) error {
if !dbConfig.EmailOneTimeAccessAsAdminEnabled.IsTrue() {
return &common.OneTimeAccessDisabledError{}
}
_, err := s.requestOneTimeAccessEmailInternal(ctx, userID, "", ttl, false, dbConfig)
return err
}
func (s *Service) RequestOneTimeAccessEmailAsUnauthenticatedUser(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID, redirectPath string) (string, error) {
if !dbConfig.EmailOneTimeAccessAsUnauthenticatedEnabled.IsTrue() {
return "", &common.OneTimeAccessDisabledError{}
}
var userId string
err := s.db.Model(&model.User{}).Select("id").Where("email = ?", userID).First(&userId).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
// Do not return error if user not found to prevent email enumeration
return "", nil
} else if err != nil {
return "", err
}
deviceToken, err := s.requestOneTimeAccessEmailInternal(ctx, userId, redirectPath, 15*time.Minute, true, dbConfig)
if err != nil {
return "", err
} else if deviceToken == nil {
return "", errors.New("device token expected but not returned")
}
return *deviceToken, nil
}
func (s *Service) requestOneTimeAccessEmailInternal(ctx context.Context, userID, redirectPath string, ttl time.Duration, withDeviceToken bool, dbConfig *appconfig.AppConfigModel) (*string, error) {
// Load the user to ensure it exists and has an email address
user, err := s.userProvider.GetUser(ctx, userID)
if err != nil {
return nil, err
}
if user.Email == nil {
return nil, &common.UserEmailNotSetError{}
}
oneTimeAccessToken, deviceToken, err := StoreToken(ctx, s.actorService, user.ID, ttl, withDeviceToken)
if err != nil {
return nil, err
}
go func() {
// This runs in background, so use a context without cancellation (or it would be stopped when the request ends)
// We still want to have a context derived from the request's to carry over tracing info
innerCtx := context.WithoutCancel(ctx)
link := common.EnvConfig.AppURL + "/lc"
linkWithCode := link + "/" + oneTimeAccessToken
// Add redirect path to the link
if strings.HasPrefix(redirectPath, "/") {
encodedRedirectPath := url.QueryEscape(redirectPath)
linkWithCode = linkWithCode + "?redirect=" + encodedRedirectPath
}
innerErr := s.emailSender.SendOneTimeAccessEmail(innerCtx, dbConfig, email.Address{
Name: user.FullName(),
Email: *user.Email,
}, EmailData{
Code: oneTimeAccessToken,
LoginLink: link,
LoginLinkWithCode: linkWithCode,
ExpirationString: utils.DurationToString(ttl),
})
if innerErr != nil {
slog.ErrorContext(innerCtx, "Failed to send one-time access token email", slog.Any("error", innerErr), slog.String("address", *user.Email))
return
}
}()
return deviceToken, nil
}
func (s *Service) CreateToken(ctx context.Context, userID string, ttl time.Duration) (token string, err error) {
// Load the user to ensure it exists
_, err = s.userProvider.GetUser(ctx, userID)
if errors.Is(err, gorm.ErrRecordNotFound) {
return "", &common.UserNotFoundError{}
} else if err != nil {
return "", err
}
token, _, err = StoreToken(ctx, s.actorService, userID, ttl, false)
if err != nil {
return "", err
}
return token, nil
}
func (s *Service) ExchangeToken(ctx context.Context, dbConfig *appconfig.AppConfigModel, token, deviceToken, ipAddress, userAgent string) (model.User, string, 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, TokenActorType, token, tokenMethodConsume, tokenConsumeRequest{
DeviceToken: deviceToken,
})
if err != nil {
return model.User{}, "", fmt.Errorf("error invoking one-time access token actor: %w", err)
}
var consumeRes tokenConsumeResponse
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 tokenConsumeNotFound:
return model.User{}, "", &common.TokenInvalidOrExpiredError{}
case tokenConsumeDeviceMismatch:
return model.User{}, "", &common.DeviceCodeInvalid{}
case tokenConsumeOK:
// 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.completeTokenExchange(ctx, dbConfig, consumeRes.State, ipAddress, userAgent)
if err != nil {
s.restoreToken(ctx, token, consumeRes.State)
return model.User{}, "", err
}
return user, accessToken, nil
}
// completeTokenExchange performs the work that follows consuming a token: loading the user, validating it, and issuing an access token.
func (s *Service) completeTokenExchange(ctx context.Context, dbConfig *appconfig.AppConfigModel, state TokenState, 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 user.Disabled {
return model.User{}, "", &common.UserDisabledError{}
}
accessToken, err := s.signer.GenerateAccessToken(
user,
authenticationMethodOneTimePassword,
dbConfig.SessionDuration.AsDurationMinutes(),
)
if err != nil {
return model.User{}, "", err
}
s.auditLog.Create(
ctx, model.AuditLogEventOneTimeAccessTokenSignIn,
ipAddress, userAgent,
user.ID,
model.AuditLogData{},
s.db,
)
return user, accessToken, nil
}
// restoreToken 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 *Service) restoreToken(parentCtx context.Context, token string, state TokenState) {
// 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, TokenActorType, token, TokenMethodRestore, state)
if err != nil {
slog.ErrorContext(ctx, "Failed to restore one-time access token after a failed exchange", slog.Any("error", err))
}
}
// StoreToken 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 StoreToken(ctx context.Context, store TokenStore, userID string, ttl time.Duration, withDeviceToken bool) (token string, deviceToken *string, err error) {
token, deviceToken, err = generateToken(ttl, withDeviceToken)
if err != nil {
return "", nil, err
}
now := time.Now().Round(time.Second)
state := TokenState{
UserID: userID,
DeviceToken: deviceToken,
ExpiresAt: now.Add(ttl),
}
err = store.SetState(ctx, TokenActorType, 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
}
// generateToken generates the random token value (and optional device token) for a one-time access token.
func generateToken(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)
if err != nil {
return "", nil, err
}
if withDeviceToken {
dt, err := utils.GenerateRandomAlphanumericString(16)
if err != nil {
return "", nil, err
}
deviceToken = &dt
}
return token, deviceToken, nil
}
@@ -0,0 +1,172 @@
package onetimeaccess
import (
"context"
"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"
"github.com/pocket-id/pocket-id/backend/internal/utils/email"
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
)
type fakeSigner struct{}
func (fakeSigner) GenerateAccessToken(_ model.User, _ string, _ time.Duration) (string, error) {
return "access-token", nil
}
type fakeAuditLogger struct {
events []model.AuditLogEvent
}
func (f *fakeAuditLogger) Create(_ context.Context, event model.AuditLogEvent, _, _, _ string, _ model.AuditLogData, _ *gorm.DB) (model.AuditLog, bool) {
f.events = append(f.events, event)
return model.AuditLog{}, true
}
type fakeUserProvider struct {
db *gorm.DB
}
func (f fakeUserProvider) GetUser(ctx context.Context, userID string) (model.User, error) {
var user model.User
err := f.db.WithContext(ctx).Where("id = ?", userID).First(&user).Error
return user, err
}
type fakeEmailSender struct{}
func (fakeEmailSender) SendOneTimeAccessEmail(_ context.Context, _ *appconfig.AppConfigModel, _ email.Address, _ EmailData) error {
return nil
}
// newServiceForTest sets up a Service backed by an in-memory test actor host, and returns it together with the host and the audit logger it records into
func newServiceForTest(t *testing.T, db *gorm.DB) (*Service, *local.Host, *fakeAuditLogger) {
t.Helper()
auditLog := &fakeAuditLogger{}
var svc *Service
host := testutils.NewActorHostForTest(t, func(t *testing.T, h *local.Host) {
err := h.RegisterActor(TokenActorType, NewTokenActor)
require.NoError(t, err)
svc = newService(Dependencies{
DB: db,
Signer: fakeSigner{},
AuditLog: auditLog,
UserProvider: fakeUserProvider{db: db},
EmailSender: fakeEmailSender{},
}, h.Service())
})
require.NotNil(t, svc)
return svc, host, auditLog
}
func TestExchangeTokenSuccess(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
svc, host, auditLog := newServiceForTest(t, db)
user := model.User{
Base: model.Base{ID: "enabled-user"},
Username: "enabled-user",
}
require.NoError(t, db.Create(&user).Error)
token, _, err := StoreToken(t.Context(), svc.actorService, user.ID, time.Minute, false)
require.NoError(t, err)
dbConfig := appconfig.NewTestConfig(nil)
exchangedUser, accessToken, err := svc.ExchangeToken(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 TokenState
err = host.GetState(t.Context(), TokenActorType, token, &state)
require.ErrorIs(t, err, actor.ErrStateNotFound)
// A sign-in audit log must have been created
require.Equal(t, []model.AuditLogEvent{model.AuditLogEventOneTimeAccessTokenSignIn}, auditLog.events)
}
func TestExchangeTokenInvalidToken(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
svc, _, _ := newServiceForTest(t, db)
dbConfig := appconfig.NewTestConfig(nil)
_, _, err := svc.ExchangeToken(t.Context(), dbConfig, "does-not-exist", "", "", "")
var invalidErr *common.TokenInvalidOrExpiredError
require.ErrorAs(t, err, &invalidErr)
}
func TestExchangeTokenDeviceMismatch(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
svc, host, _ := newServiceForTest(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 := StoreToken(t.Context(), svc.actorService, user.ID, time.Minute, true)
require.NoError(t, err)
require.NotNil(t, deviceToken)
dbConfig := appconfig.NewTestConfig(nil)
_, _, err = svc.ExchangeToken(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 TokenState
err = host.GetState(t.Context(), TokenActorType, token, &state)
require.NoError(t, err)
require.Equal(t, user.ID, state.UserID)
}
func TestExchangeTokenRejectsDisabledUser(t *testing.T) {
db := testutils.NewDatabaseForTest(t)
svc, host, auditLog := newServiceForTest(t, db)
user := model.User{
Base: model.Base{ID: "disabled-user"},
Username: "disabled-user",
Disabled: true,
}
require.NoError(t, db.Create(&user).Error)
// Store a one-time access token for the disabled user in the actor state store
token, _, err := StoreToken(t.Context(), svc.actorService, user.ID, time.Minute, false)
require.NoError(t, err)
dbConfig := appconfig.NewTestConfig(nil)
exchangedUser, accessToken, err := svc.ExchangeToken(t.Context(), dbConfig, token, "", "", "")
var userDisabledErr *common.UserDisabledError
require.ErrorAs(t, err, &userDisabledErr)
require.Empty(t, exchangedUser.ID)
require.Empty(t, accessToken)
// The token must have been restored (not consumed), since the exchange failed because the user is disabled
var state TokenState
err = host.GetState(t.Context(), TokenActorType, token, &state)
require.NoError(t, err)
require.Equal(t, user.ID, state.UserID)
require.Empty(t, auditLog.events)
}