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