mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-01 23:39:05 +02:00
refactor: use actors for db configuration (#1604)
Co-authored-by: Claude <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude
parent
472fff33ea
commit
2cfbcb4b67
@@ -1,6 +1,7 @@
|
||||
package webauthn
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -13,16 +14,22 @@ import (
|
||||
|
||||
type handler struct {
|
||||
service *Service
|
||||
appConfig AppConfigProvider
|
||||
appConfig AppConfigResolver
|
||||
}
|
||||
|
||||
func newHandler(service *Service, appConfig AppConfigProvider) *handler {
|
||||
func newHandler(service *Service, appConfig AppConfigResolver) *handler {
|
||||
return &handler{service: service, appConfig: appConfig}
|
||||
}
|
||||
|
||||
func (h *handler) beginRegistration(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
|
||||
}
|
||||
|
||||
userID := c.GetString("userID")
|
||||
options, err := h.service.BeginRegistration(c.Request.Context(), userID)
|
||||
options, err := h.service.BeginRegistration(c.Request.Context(), dbConfig, userID)
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
@@ -67,6 +74,12 @@ func (h *handler) beginLogin(c *gin.Context) {
|
||||
}
|
||||
|
||||
func (h *handler) verifyLogin(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
|
||||
}
|
||||
|
||||
sessionID, err := c.Cookie(cookie.SessionIdCookieName)
|
||||
if err != nil {
|
||||
_ = c.Error(&common.MissingSessionIdError{})
|
||||
@@ -79,7 +92,7 @@ func (h *handler) verifyLogin(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
user, token, err := h.service.VerifyLogin(c.Request.Context(), sessionID, credentialAssertionData, c.ClientIP(), c.Request.UserAgent())
|
||||
user, token, err := h.service.VerifyLogin(c.Request.Context(), dbConfig, sessionID, credentialAssertionData, c.ClientIP(), c.Request.UserAgent())
|
||||
if err != nil {
|
||||
_ = c.Error(err)
|
||||
return
|
||||
@@ -91,7 +104,7 @@ func (h *handler) verifyLogin(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
maxAge := int(h.appConfig.GetDbConfig().SessionDuration.AsDurationMinutes().Seconds())
|
||||
maxAge := int(dbConfig.SessionDuration.AsDurationMinutes().Seconds())
|
||||
cookie.AddAccessTokenCookie(c, maxAge, token)
|
||||
|
||||
c.JSON(http.StatusOK, userDto)
|
||||
|
||||
@@ -8,22 +8,24 @@ import (
|
||||
"github.com/lestrrat-go/jwx/v3/jwt"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
)
|
||||
|
||||
type TokenService interface {
|
||||
GenerateAccessToken(user model.User, authenticationMethod string) (string, error)
|
||||
GenerateAccessToken(user model.User, authenticationMethod string, sessionDuration time.Duration) (string, error)
|
||||
VerifyAccessToken(tokenString string) (jwt.Token, error)
|
||||
GetAuthenticationMethod(token jwt.Token) (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)
|
||||
CreateNewSignInWithEmail(ctx context.Context, ipAddress, userAgent, userID string, tx *gorm.DB) model.AuditLog
|
||||
CreateNewSignInWithEmail(ctx context.Context, ipAddress, userAgent, userID string, tx *gorm.DB, emailLoginNotificationEnabled bool) model.AuditLog
|
||||
}
|
||||
|
||||
type AppConfigProvider interface {
|
||||
GetDbConfig() *model.AppConfig
|
||||
// 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 {
|
||||
@@ -32,7 +34,7 @@ type Dependencies struct {
|
||||
|
||||
Signer TokenService
|
||||
AuditLog AuditLogger
|
||||
AppConfig AppConfigProvider
|
||||
AppConfig AppConfigResolver
|
||||
}
|
||||
|
||||
type Module struct {
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"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"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
@@ -23,17 +24,19 @@ import (
|
||||
// It must match the value emitted by the JWT service in the access token's "amr" claim
|
||||
const authenticationMethodPhishingResistant = "phr"
|
||||
|
||||
const defaultRPDisplayName = "Pocket ID"
|
||||
|
||||
type Service struct {
|
||||
db *gorm.DB
|
||||
webAuthn *gowebauthn.WebAuthn
|
||||
signer TokenService
|
||||
auditLog AuditLogger
|
||||
appConfig AppConfigProvider
|
||||
db *gorm.DB
|
||||
webAuthn *gowebauthn.WebAuthn
|
||||
signer TokenService
|
||||
auditLog AuditLogger
|
||||
}
|
||||
|
||||
func newService(deps Dependencies) (*Service, error) {
|
||||
wa, err := gowebauthn.New(&gowebauthn.Config{
|
||||
RPDisplayName: deps.AppConfig.GetDbConfig().AppName.Value,
|
||||
// Set a default value, it will be set again later
|
||||
RPDisplayName: defaultRPDisplayName,
|
||||
RPID: utils.GetHostnameFromURL(deps.AppURL),
|
||||
RPOrigins: []string{deps.AppURL},
|
||||
AuthenticatorSelection: protocol.AuthenticatorSelection{
|
||||
@@ -57,22 +60,21 @@ func newService(deps Dependencies) (*Service, error) {
|
||||
}
|
||||
|
||||
return &Service{
|
||||
db: deps.DB,
|
||||
webAuthn: wa,
|
||||
signer: deps.Signer,
|
||||
auditLog: deps.AuditLog,
|
||||
appConfig: deps.AppConfig,
|
||||
db: deps.DB,
|
||||
webAuthn: wa,
|
||||
signer: deps.Signer,
|
||||
auditLog: deps.AuditLog,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) BeginRegistration(ctx context.Context, userID string) (*PublicKeyCredentialCreationOptions, error) {
|
||||
func (s *Service) BeginRegistration(ctx context.Context, dbConfig *appconfig.AppConfigModel, userID string) (*PublicKeyCredentialCreationOptions, error) {
|
||||
s.updateWebAuthnConfig(dbConfig)
|
||||
|
||||
tx := s.db.Begin()
|
||||
defer func() {
|
||||
tx.Rollback()
|
||||
}()
|
||||
|
||||
s.updateWebAuthnConfig()
|
||||
|
||||
var user model.User
|
||||
err := tx.
|
||||
WithContext(ctx).
|
||||
@@ -227,7 +229,7 @@ func (s *Service) BeginLogin(ctx context.Context) (*PublicKeyCredentialRequestOp
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) VerifyLogin(ctx context.Context, sessionID string, credentialAssertionData *protocol.ParsedCredentialAssertionData, ipAddress, userAgent string) (model.User, string, error) {
|
||||
func (s *Service) VerifyLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, sessionID string, credentialAssertionData *protocol.ParsedCredentialAssertionData, ipAddress, userAgent string) (model.User, string, error) {
|
||||
tx := s.db.Begin()
|
||||
defer func() {
|
||||
tx.Rollback()
|
||||
@@ -270,12 +272,12 @@ func (s *Service) VerifyLogin(ctx context.Context, sessionID string, credentialA
|
||||
return model.User{}, "", &common.UserDisabledError{}
|
||||
}
|
||||
|
||||
token, err := s.signer.GenerateAccessToken(*user, authenticationMethodPhishingResistant)
|
||||
token, err := s.signer.GenerateAccessToken(*user, authenticationMethodPhishingResistant, dbConfig.SessionDuration.AsDurationMinutes())
|
||||
if err != nil {
|
||||
return model.User{}, "", err
|
||||
}
|
||||
|
||||
s.auditLog.CreateNewSignInWithEmail(ctx, ipAddress, userAgent, user.ID, tx)
|
||||
s.auditLog.CreateNewSignInWithEmail(ctx, ipAddress, userAgent, user.ID, tx, dbConfig.EmailLoginNotificationEnabled.IsTrue())
|
||||
|
||||
err = tx.Commit().Error
|
||||
if err != nil {
|
||||
@@ -373,8 +375,8 @@ func (s *Service) UpdateCredential(ctx context.Context, userID, credentialID, na
|
||||
}
|
||||
|
||||
// updateWebAuthnConfig updates the WebAuthn configuration with the app name as it can change during runtime
|
||||
func (s *Service) updateWebAuthnConfig() {
|
||||
s.webAuthn.Config.RPDisplayName = s.appConfig.GetDbConfig().AppName.Value
|
||||
func (s *Service) updateWebAuthnConfig(dbConfig *appconfig.AppConfigModel) {
|
||||
s.webAuthn.Config.RPDisplayName = dbConfig.AppName.String()
|
||||
}
|
||||
|
||||
func (s *Service) CreateReauthenticationTokenWithAccessToken(ctx context.Context, accessToken string) (string, error) {
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"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"
|
||||
@@ -28,7 +29,7 @@ func newFakeSigner() *fakeSigner {
|
||||
return &fakeSigner{tokens: map[string]jwt.Token{}}
|
||||
}
|
||||
|
||||
func (s *fakeSigner) GenerateAccessToken(user model.User, authenticationMethod string) (string, error) {
|
||||
func (s *fakeSigner) GenerateAccessToken(user model.User, authenticationMethod string, _ time.Duration) (string, error) {
|
||||
builder := jwt.NewBuilder().
|
||||
Subject(user.ID).
|
||||
IssuedAt(time.Now())
|
||||
@@ -85,7 +86,7 @@ func TestCreateReauthenticationTokenWithAccessToken(t *testing.T) {
|
||||
|
||||
t.Run("accepts a fresh access token from WebAuthn login", func(t *testing.T) {
|
||||
service, signer, user := setupService(t)
|
||||
accessToken, err := signer.GenerateAccessToken(user, authenticationMethodPhishingResistant)
|
||||
accessToken, err := signer.GenerateAccessToken(user, authenticationMethodPhishingResistant, time.Hour)
|
||||
require.NoError(t, err)
|
||||
|
||||
reauthenticationToken, err := service.CreateReauthenticationTokenWithAccessToken(t.Context(), accessToken)
|
||||
@@ -96,7 +97,7 @@ func TestCreateReauthenticationTokenWithAccessToken(t *testing.T) {
|
||||
|
||||
t.Run("rejects a fresh access token from one-time access login", func(t *testing.T) {
|
||||
service, signer, user := setupService(t)
|
||||
accessToken, err := signer.GenerateAccessToken(user, "otp")
|
||||
accessToken, err := signer.GenerateAccessToken(user, "otp", time.Hour)
|
||||
require.NoError(t, err)
|
||||
|
||||
reauthenticationToken, err := service.CreateReauthenticationTokenWithAccessToken(t.Context(), accessToken)
|
||||
@@ -108,7 +109,7 @@ func TestCreateReauthenticationTokenWithAccessToken(t *testing.T) {
|
||||
|
||||
t.Run("rejects a fresh access token without an authentication method", func(t *testing.T) {
|
||||
service, signer, user := setupService(t)
|
||||
accessToken, err := signer.GenerateAccessToken(user, "")
|
||||
accessToken, err := signer.GenerateAccessToken(user, "", time.Hour)
|
||||
require.NoError(t, err)
|
||||
|
||||
reauthenticationToken, err := service.CreateReauthenticationTokenWithAccessToken(t.Context(), accessToken)
|
||||
@@ -119,6 +120,18 @@ func TestCreateReauthenticationTokenWithAccessToken(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestWebAuthnDisplayNameUsesRequestConfig(t *testing.T) {
|
||||
service, err := newService(Dependencies{
|
||||
DB: testutils.NewDatabaseForTest(t),
|
||||
AppURL: "https://example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, defaultRPDisplayName, service.webAuthn.Config.RPDisplayName)
|
||||
|
||||
service.updateWebAuthnConfig(&appconfig.AppConfigModel{AppName: "Custom App"})
|
||||
require.Equal(t, "Custom App", service.webAuthn.Config.RPDisplayName)
|
||||
}
|
||||
|
||||
func TestConsumeReauthenticationTokenReturnsTokenCreationTime(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
service := &Service{db: db}
|
||||
|
||||
Reference in New Issue
Block a user