mirror of
https://github.com/pocket-id/pocket-id.git
synced 2026-10-08 10:49:06 +02:00
feat: add configurable login notification modes (#1813)
This commit is contained in:
@@ -0,0 +1,33 @@
|
||||
package appconfig
|
||||
|
||||
import "encoding/json"
|
||||
|
||||
const (
|
||||
LoginNotificationDisabled = "disabled"
|
||||
LoginNotificationAlways = "always"
|
||||
LoginNotificationIPAndUserAgent = "ipAndUserAgent"
|
||||
LoginNotificationBrowserRecognition = "browserRecognition"
|
||||
)
|
||||
|
||||
// UnmarshalJSON preserves the notification policy when loading state written before notification modes existed
|
||||
func (m *AppConfigModel) UnmarshalJSON(data []byte) error {
|
||||
type config AppConfigModel
|
||||
value := struct {
|
||||
*config
|
||||
LegacyNotificationEnabled AppConfigValue `json:"emailLoginNotificationEnabled"`
|
||||
}{config: (*config)(m)}
|
||||
if err := json.Unmarshal(data, &value); err != nil {
|
||||
return err
|
||||
}
|
||||
if m.EmailLoginNotificationMode == "" && value.LegacyNotificationEnabled != "" {
|
||||
m.EmailLoginNotificationMode = legacyLoginNotificationMode(value.LegacyNotificationEnabled)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func legacyLoginNotificationMode(enabled AppConfigValue) AppConfigValue {
|
||||
if enabled.IsTrue() {
|
||||
return LoginNotificationBrowserRecognition
|
||||
}
|
||||
return LoginNotificationDisabled
|
||||
}
|
||||
@@ -59,6 +59,9 @@ func LoadLegacyConfig(ctx context.Context, db *gorm.DB) (map[string]string, erro
|
||||
func fromLegacyConfig(legacyCfg map[string]string) (*AppConfigModel, error) {
|
||||
// Start from the default configuration, then override with the values from the legacy config
|
||||
dest := getDefaultConfig()
|
||||
if legacyCfg["emailLoginNotificationMode"] == "" {
|
||||
dest.EmailLoginNotificationMode = legacyLoginNotificationMode(AppConfigValue(legacyCfg["emailLoginNotificationEnabled"]))
|
||||
}
|
||||
|
||||
rt := reflect.ValueOf(dest).Elem().Type()
|
||||
rv := reflect.ValueOf(dest).Elem()
|
||||
|
||||
@@ -35,7 +35,7 @@ type AppConfigModel struct {
|
||||
SmtpPassword AppConfigValue `json:"smtpPassword" env:"SMTP_PASSWORD" sensitive:"true"`
|
||||
SmtpTls AppConfigValue `json:"smtpTls" env:"SMTP_TLS"`
|
||||
SmtpSkipCertVerify AppConfigValue `json:"smtpSkipCertVerify" env:"SMTP_SKIP_CERT_VERIFY" type:"bool"`
|
||||
EmailLoginNotificationEnabled AppConfigValue `json:"emailLoginNotificationEnabled" env:"EMAIL_LOGIN_NOTIFICATION_ENABLED" type:"bool"`
|
||||
EmailLoginNotificationMode AppConfigValue `json:"emailLoginNotificationMode" env:"EMAIL_LOGIN_NOTIFICATION_MODE"`
|
||||
EmailOneTimeAccessAsUnauthenticatedEnabled AppConfigValue `json:"emailOneTimeAccessAsUnauthenticatedEnabled" env:"EMAIL_ONE_TIME_ACCESS_AS_UNAUTHENTICATED_ENABLED" type:"bool" public:"true"`
|
||||
EmailOneTimeAccessAsAdminEnabled AppConfigValue `json:"emailOneTimeAccessAsAdminEnabled" env:"EMAIL_ONE_TIME_ACCESS_AS_ADMIN_ENABLED" type:"bool" public:"true"`
|
||||
EmailApiKeyExpirationEnabled AppConfigValue `json:"emailApiKeyExpirationEnabled" env:"EMAIL_API_KEY_EXPIRATION_ENABLED" type:"bool"`
|
||||
@@ -133,15 +133,15 @@ func getDefaultConfig() *AppConfigModel {
|
||||
SignupDefaultCustomClaims: "[]",
|
||||
AccentColor: "default",
|
||||
// Email
|
||||
RequireUserEmail: "true",
|
||||
SmtpHost: "",
|
||||
SmtpPort: "",
|
||||
SmtpFrom: "",
|
||||
SmtpUser: "",
|
||||
SmtpPassword: "",
|
||||
SmtpTls: "none",
|
||||
SmtpSkipCertVerify: "false",
|
||||
EmailLoginNotificationEnabled: "false",
|
||||
RequireUserEmail: "true",
|
||||
SmtpHost: "",
|
||||
SmtpPort: "",
|
||||
SmtpFrom: "",
|
||||
SmtpUser: "",
|
||||
SmtpPassword: "",
|
||||
SmtpTls: "none",
|
||||
SmtpSkipCertVerify: "false",
|
||||
EmailLoginNotificationMode: LoginNotificationDisabled,
|
||||
EmailOneTimeAccessAsUnauthenticatedEnabled: "false",
|
||||
EmailOneTimeAccessAsAdminEnabled: "false",
|
||||
EmailApiKeyExpirationEnabled: "false",
|
||||
|
||||
@@ -236,6 +236,17 @@ func (s *AppConfigService) loadDbConfigFromEnv() (*AppConfigModel, error) {
|
||||
}
|
||||
}
|
||||
|
||||
// TODO: Remove in next major version (v3)
|
||||
// Preserve the old environment setting unless an explicit notification mode replaces it
|
||||
if _, ok := os.LookupEnv("EMAIL_LOGIN_NOTIFICATION_MODE"); !ok {
|
||||
if enabled, exists := os.LookupEnv("EMAIL_LOGIN_NOTIFICATION_ENABLED"); exists {
|
||||
if enabled != "true" && enabled != "false" {
|
||||
return nil, errors.New("EMAIL_LOGIN_NOTIFICATION_ENABLED must be true or false")
|
||||
}
|
||||
dest.EmailLoginNotificationMode = legacyLoginNotificationMode(AppConfigValue(enabled))
|
||||
}
|
||||
}
|
||||
|
||||
// Validate the resolved configuration before exposing values to the rest of the application
|
||||
err := validateEnvConfig(dest)
|
||||
if err != nil {
|
||||
|
||||
@@ -9,52 +9,58 @@ import (
|
||||
"github.com/italypaleale/francis/builtin/cronjob"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
)
|
||||
|
||||
const (
|
||||
// cleanupJobInterval is how often the audit log cleanup job runs
|
||||
// cleanupJobInterval is how often each audit cleanup job runs
|
||||
cleanupJobInterval = 24 * time.Hour
|
||||
// cleanupJobJitter spreads each occurrence around its scheduled time, so the cleanup jobs don't all hit the database at once
|
||||
cleanupJobJitter = 5 * time.Minute
|
||||
)
|
||||
|
||||
type cleanupJob struct {
|
||||
type cleanupJobs struct {
|
||||
db *gorm.DB
|
||||
retentionDays int
|
||||
}
|
||||
|
||||
// newCleanupJob returns the cron job actor that deletes audit logs past the retention window
|
||||
func newCleanupJob(db *gorm.DB, retentionDays int) (*cronjob.CronJob, error) {
|
||||
job := &cleanupJob{
|
||||
db: db,
|
||||
retentionDays: retentionDays,
|
||||
// newCleanupJobs returns the cron job actors for audit retention
|
||||
func newCleanupJobs(db *gorm.DB, retentionDays int) ([]*cronjob.CronJob, error) {
|
||||
jobs := &cleanupJobs{db: db, retentionDays: retentionDays}
|
||||
|
||||
clearAuditLogs, err := newCleanupJob("ClearAuditLogs", jobs.clearAuditLogs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return []*cronjob.CronJob{clearAuditLogs}, nil
|
||||
}
|
||||
|
||||
// newCleanupJob applies the shared daily schedule to each audit cleanup
|
||||
func newCleanupJob(name string, fn func(context.Context) error) (*cronjob.CronJob, error) {
|
||||
cronActor, err := cronjob.New(
|
||||
"ClearAuditLogs",
|
||||
cronjob.WithJob(job.clearAuditLogs),
|
||||
name,
|
||||
cronjob.WithJob(fn),
|
||||
cronjob.WithInterval(cleanupJobInterval),
|
||||
cronjob.WithJitter(cleanupJobJitter),
|
||||
// Also run right after the job is first registered, so rows that aged out while Pocket ID wasn't running are removed at startup
|
||||
// Also run right after the job is first registered, so rows that expired while Pocket ID wasn't running are removed at startup
|
||||
cronjob.WithImmediate(),
|
||||
cronjob.WithLogger(slog.Default()),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error creating audit log cleanup cron job: %w", err)
|
||||
return nil, fmt.Errorf("error creating %s cron job: %w", name, err)
|
||||
}
|
||||
|
||||
return cronActor, nil
|
||||
}
|
||||
|
||||
// clearAuditLogs deletes audit logs older than the configured retention window
|
||||
func (j *cleanupJob) clearAuditLogs(ctx context.Context) error {
|
||||
func (j *cleanupJobs) clearAuditLogs(ctx context.Context) error {
|
||||
cutoff := time.Now().AddDate(0, 0, -j.retentionDays)
|
||||
|
||||
st := j.db.
|
||||
WithContext(ctx).
|
||||
Delete(&model.AuditLog{}, "created_at < ?", datatype.DateTime(cutoff))
|
||||
Delete(&AuditLog{}, "created_at < ?", datatype.DateTime(cutoff))
|
||||
if st.Error != nil {
|
||||
return fmt.Errorf("failed to delete old audit logs: %w", st.Error)
|
||||
}
|
||||
|
||||
@@ -12,69 +12,35 @@ import (
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
)
|
||||
|
||||
func TestModuleRegistersAuditLogCleanupCronJob(t *testing.T) {
|
||||
func TestModuleCleanupDeletesLogsPastRetention(t *testing.T) {
|
||||
const retentionDays = 7
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
|
||||
// Preserve a recent record while proving that the registered job honors the configured retention window
|
||||
require.NoError(t, db.Create(&AuditLog{Base: model.Base{ID: "expired-log"}, Event: EventSignIn}).Error)
|
||||
require.NoError(t, db.Create(&AuditLog{Base: model.Base{ID: "recent-log"}, Event: EventSignIn}).Error)
|
||||
require.NoError(t, db.Model(&AuditLog{}).Where("id = ?", "expired-log").Update("created_at", datatype.DateTime(time.Now().AddDate(0, 0, -retentionDays-1))).Error)
|
||||
|
||||
testutils.NewActorHostForTest(t, func(t *testing.T, host *local.Host) {
|
||||
t.Helper()
|
||||
_, err := New(Dependencies{
|
||||
DB: db,
|
||||
Actors: host,
|
||||
RetentionDays: 90,
|
||||
DB: db, Actors: host, RetentionDays: retentionDays,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
require.Eventually(t, func() bool {
|
||||
var remaining []string
|
||||
if db.Model(&AuditLog{}).Pluck("id", &remaining).Error != nil {
|
||||
return false
|
||||
}
|
||||
return len(remaining) == 1 && remaining[0] == "recent-log"
|
||||
}, 5*time.Second, 10*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestModuleRequiresActorHostForCleanupJob(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
|
||||
_, err := New(Dependencies{DB: db, RetentionDays: 90})
|
||||
require.ErrorContains(t, err, "actor host is required")
|
||||
|
||||
// With the cleanup disabled there is nothing to register, so the actor host is not needed
|
||||
_, err = New(Dependencies{DB: db, RetentionDays: 90, CleanupDisabled: true})
|
||||
func TestNewCleanupJobsPreserveAuditActorName(t *testing.T) {
|
||||
jobs, err := newCleanupJobs(testutils.NewDatabaseForTest(t), 90)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func TestAuditLogCleanupJobDeletesLogsPastRetention(t *testing.T) {
|
||||
const retentionDays = 90
|
||||
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
user := model.User{
|
||||
Base: model.Base{ID: "cleanup-job-user"},
|
||||
Username: "cleanup-job-user",
|
||||
FirstName: "Cleanup",
|
||||
LastName: "Job",
|
||||
DisplayName: "Cleanup Job",
|
||||
}
|
||||
err := db.Create(&user).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
err = db.Create(&model.AuditLog{Base: model.Base{ID: "log-old"}, Event: model.AuditLogEventSignIn, UserID: user.ID}).Error
|
||||
require.NoError(t, err)
|
||||
err = db.Create(&model.AuditLog{Base: model.Base{ID: "log-recent"}, Event: model.AuditLogEventSignIn, UserID: user.ID}).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
// BeforeCreate stamps CreatedAt, so the log past the retention window is backdated directly
|
||||
oldCreatedAt := datatype.DateTime(time.Now().AddDate(0, 0, -retentionDays-1))
|
||||
err = db.Model(&model.AuditLog{}).Where("id = ?", "log-old").Update("created_at", oldCreatedAt).Error
|
||||
require.NoError(t, err)
|
||||
|
||||
job := &cleanupJob{db: db, retentionDays: retentionDays}
|
||||
err = job.clearAuditLogs(t.Context())
|
||||
require.NoError(t, err)
|
||||
|
||||
var remaining []string
|
||||
err = db.Model(&model.AuditLog{}).Pluck("id", &remaining).Error
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"log-recent"}, remaining)
|
||||
}
|
||||
|
||||
func TestNewCleanupJobCreatesCronActor(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
|
||||
cronActor, err := newCleanupJob(db, 90)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "cronjob.ClearAuditLogs", cronActor.ActorType())
|
||||
require.Len(t, jobs, 1)
|
||||
require.Equal(t, "cronjob.ClearAuditLogs", jobs[0].ActorType())
|
||||
}
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
package dto
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
)
|
||||
|
||||
type AuditLogDto struct {
|
||||
type auditLogDto struct {
|
||||
ID string `json:"id"`
|
||||
CreatedAt datatype.DateTime `json:"createdAt"`
|
||||
|
||||
+24
-38
@@ -1,34 +1,20 @@
|
||||
package controller
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/middleware"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/service"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
)
|
||||
|
||||
// NewAuditLogController creates a new controller for audit log management
|
||||
// @Summary Audit log controller
|
||||
// @Description Initializes API endpoints for accessing audit logs
|
||||
// @Tags Audit Logs
|
||||
func NewAuditLogController(group *gin.RouterGroup, auditLogService *service.AuditLogService, authMiddleware *middleware.AuthMiddleware) {
|
||||
alc := AuditLogController{
|
||||
auditLogService: auditLogService,
|
||||
}
|
||||
|
||||
group.GET("/audit-logs/all", authMiddleware.Add(), httpserver.Handle(alc.listAllAuditLogsHandler))
|
||||
group.GET("/audit-logs", authMiddleware.WithAdminNotRequired().Add(), httpserver.Handle(alc.listAuditLogsForUserHandler))
|
||||
group.GET("/audit-logs/filters/client-names", authMiddleware.Add(), httpserver.Handle(alc.listClientNamesHandler))
|
||||
group.GET("/audit-logs/filters/users", authMiddleware.Add(), httpserver.Handle(alc.listUserNamesWithIdsHandler))
|
||||
type handler struct {
|
||||
service *service
|
||||
}
|
||||
|
||||
type AuditLogController struct {
|
||||
auditLogService *service.AuditLogService
|
||||
func newHandler(service *service) *handler {
|
||||
return &handler{service: service}
|
||||
}
|
||||
|
||||
// listAuditLogsForUserHandler godoc
|
||||
@@ -39,22 +25,22 @@ type AuditLogController struct {
|
||||
// @Param pagination[limit] query int false "Number of items per page" default(20)
|
||||
// @Param sort[column] query string false "Column to sort by"
|
||||
// @Param sort[direction] query string false "Sort direction (asc or desc)" default("asc")
|
||||
// @Success 200 {object} dto.Paginated[dto.AuditLogDto]
|
||||
// @Success 200 {object} dto.Paginated[auditLogDto]
|
||||
// @Failure default {object} dto.ErrorDto "Error"
|
||||
// @Router /api/audit-logs [get]
|
||||
func (alc *AuditLogController) listAuditLogsForUserHandler(c *gin.Context) error {
|
||||
func (h *handler) listAuditLogsForUserHandler(c *gin.Context) error {
|
||||
listRequestOptions := utils.ParseListRequestOptions(c)
|
||||
|
||||
userID := c.GetString("userID")
|
||||
|
||||
// Fetch audit logs for the user
|
||||
logs, pagination, err := alc.auditLogService.ListAuditLogsForUser(c.Request.Context(), userID, listRequestOptions)
|
||||
logs, pagination, err := h.service.ListAuditLogsForUser(c.Request.Context(), userID, listRequestOptions)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Map the audit logs to DTOs
|
||||
var logsDtos []dto.AuditLogDto
|
||||
var logsDtos []auditLogDto
|
||||
err = dto.MapStructList(logs, &logsDtos)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -62,12 +48,12 @@ func (alc *AuditLogController) listAuditLogsForUserHandler(c *gin.Context) error
|
||||
|
||||
// Add device information to the logs
|
||||
for i, logsDto := range logsDtos {
|
||||
logsDto.Device = alc.auditLogService.DeviceStringFromUserAgent(logs[i].UserAgent)
|
||||
logsDto.Device = h.service.DeviceStringFromUserAgent(logs[i].UserAgent)
|
||||
logsDto.ActorUsername = logsDto.Data["actorUsername"]
|
||||
logsDtos[i] = logsDto
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, dto.Paginated[dto.AuditLogDto]{
|
||||
c.JSON(http.StatusOK, dto.Paginated[auditLogDto]{
|
||||
Data: logsDtos,
|
||||
Pagination: pagination,
|
||||
})
|
||||
@@ -82,31 +68,31 @@ func (alc *AuditLogController) listAuditLogsForUserHandler(c *gin.Context) error
|
||||
// @Param pagination[limit] query int false "Number of items per page" default(20)
|
||||
// @Param sort[column] query string false "Column to sort by"
|
||||
// @Param sort[direction] query string false "Sort direction (asc or desc)" default("asc")
|
||||
// @Success 200 {object} dto.Paginated[dto.AuditLogDto]
|
||||
// @Success 200 {object} dto.Paginated[auditLogDto]
|
||||
// @Failure default {object} dto.ErrorDto "Error"
|
||||
// @Router /api/audit-logs/all [get]
|
||||
func (alc *AuditLogController) listAllAuditLogsHandler(c *gin.Context) error {
|
||||
func (h *handler) listAllAuditLogsHandler(c *gin.Context) error {
|
||||
listRequestOptions := utils.ParseListRequestOptions(c)
|
||||
|
||||
logs, pagination, err := alc.auditLogService.ListAllAuditLogs(c.Request.Context(), listRequestOptions)
|
||||
logs, pagination, err := h.service.ListAllAuditLogs(c.Request.Context(), listRequestOptions)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var logsDtos []dto.AuditLogDto
|
||||
var logsDtos []auditLogDto
|
||||
err = dto.MapStructList(logs, &logsDtos)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for i, logsDto := range logsDtos {
|
||||
logsDto.Device = alc.auditLogService.DeviceStringFromUserAgent(logs[i].UserAgent)
|
||||
logsDto.Device = h.service.DeviceStringFromUserAgent(logs[i].UserAgent)
|
||||
logsDto.Username = logs[i].User.Username
|
||||
logsDto.ActorUsername = logsDto.Data["actorUsername"]
|
||||
logsDtos[i] = logsDto
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, dto.Paginated[dto.AuditLogDto]{
|
||||
c.JSON(http.StatusOK, dto.Paginated[auditLogDto]{
|
||||
Data: logsDtos,
|
||||
Pagination: pagination,
|
||||
})
|
||||
@@ -120,8 +106,8 @@ func (alc *AuditLogController) listAllAuditLogsHandler(c *gin.Context) error {
|
||||
// @Success 200 {array} string "List of client names"
|
||||
// @Failure default {object} dto.ErrorDto "Error"
|
||||
// @Router /api/audit-logs/filters/client-names [get]
|
||||
func (alc *AuditLogController) listClientNamesHandler(c *gin.Context) error {
|
||||
names, err := alc.auditLogService.ListClientNames(c.Request.Context())
|
||||
func (h *handler) listClientNamesHandler(c *gin.Context) error {
|
||||
names, err := h.service.ListClientNames(c.Request.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -137,8 +123,8 @@ func (alc *AuditLogController) listClientNamesHandler(c *gin.Context) error {
|
||||
// @Success 200 {object} map[string]string "Map of user IDs to usernames"
|
||||
// @Failure default {object} dto.ErrorDto "Error"
|
||||
// @Router /api/audit-logs/filters/users [get]
|
||||
func (alc *AuditLogController) listUserNamesWithIdsHandler(c *gin.Context) error {
|
||||
users, err := alc.auditLogService.ListUsernamesWithIds(c.Request.Context())
|
||||
func (h *handler) listUserNamesWithIdsHandler(c *gin.Context) error {
|
||||
users, err := h.service.ListUsernamesWithIds(c.Request.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,117 @@
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
)
|
||||
|
||||
func TestAuditLogRoutesPreservePermissionsAndResponses(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
module, err := New(Dependencies{DB: db, CleanupDisabled: true})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Different owners make accidental loss of the current-user filter visible
|
||||
for _, id := range []string{"alice", "bob"} {
|
||||
require.NoError(t, db.Create(&model.User{Base: model.Base{ID: id}, Username: id}).Error)
|
||||
require.NoError(t, db.Create(&AuditLog{
|
||||
Base: model.Base{ID: id + "-log"}, UserID: id, Event: EventSignIn,
|
||||
IpAddress: new("192.0.2.1"), UserAgent: "Firefox", Country: "Switzerland", City: "Zurich",
|
||||
Data: Data{"clientName": id + "-client", "actorUsername": "administrator"},
|
||||
}).Error)
|
||||
}
|
||||
|
||||
// Keep authentication lightweight while exercising which middleware each route receives
|
||||
router := gin.New()
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Next()
|
||||
if len(c.Errors) > 0 {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": c.Errors.String()})
|
||||
}
|
||||
})
|
||||
module.RegisterRoutes(router.Group("/api"), auditLogTestAuth(true), auditLogTestAuth(false))
|
||||
|
||||
for _, path := range []string{"/audit-logs", "/audit-logs/all", "/audit-logs/filters/client-names", "/audit-logs/filters/users"} {
|
||||
for _, role := range []string{"", "user", "admin"} {
|
||||
t.Run(path+"/"+role, func(t *testing.T) {
|
||||
request := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/api"+path, nil)
|
||||
request.Header.Set("X-Test-Role", role)
|
||||
response := httptest.NewRecorder()
|
||||
router.ServeHTTP(response, request)
|
||||
status := http.StatusOK
|
||||
if role == "" {
|
||||
status = http.StatusUnauthorized
|
||||
} else if role != "admin" && path != "/audit-logs" {
|
||||
status = http.StatusForbidden
|
||||
}
|
||||
require.Equal(t, status, response.Code, response.Body.String())
|
||||
if status != http.StatusOK {
|
||||
return
|
||||
}
|
||||
|
||||
assertAuditLogRouteResponse(t, path, response)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func auditLogTestAuth(adminRequired bool) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
role := c.GetHeader("X-Test-Role")
|
||||
if role == "" {
|
||||
c.AbortWithStatus(http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
if adminRequired && role != "admin" {
|
||||
c.AbortWithStatus(http.StatusForbidden)
|
||||
return
|
||||
}
|
||||
c.Set("userID", "alice")
|
||||
}
|
||||
}
|
||||
|
||||
func assertAuditLogRouteResponse(t *testing.T, path string, response *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
switch path {
|
||||
case "/audit-logs", "/audit-logs/all":
|
||||
var body struct {
|
||||
Data []map[string]any `json:"data"`
|
||||
Pagination map[string]any `json:"pagination"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal(response.Body.Bytes(), &body))
|
||||
require.NotEmpty(t, body.Pagination)
|
||||
count := 1
|
||||
if path == "/audit-logs/all" {
|
||||
count = 2
|
||||
}
|
||||
require.Len(t, body.Data, count)
|
||||
for _, entry := range body.Data {
|
||||
require.Equal(t, "SIGN_IN", entry["event"])
|
||||
require.Equal(t, "192.0.2.1", entry["ipAddress"])
|
||||
require.Equal(t, "administrator", entry["actorUsername"])
|
||||
require.Contains(t, entry, "device")
|
||||
require.Contains(t, entry, "createdAt")
|
||||
require.NotContains(t, entry, "userAgent")
|
||||
if path == "/audit-logs" {
|
||||
require.Equal(t, "alice", entry["userID"])
|
||||
} else {
|
||||
require.Equal(t, entry["userID"], entry["username"])
|
||||
}
|
||||
}
|
||||
case "/audit-logs/filters/client-names":
|
||||
var names []string
|
||||
require.NoError(t, json.Unmarshal(response.Body.Bytes(), &names))
|
||||
require.ElementsMatch(t, []string{"alice-client", "bob-client"}, names)
|
||||
case "/audit-logs/filters/users":
|
||||
var users map[string]string
|
||||
require.NoError(t, json.Unmarshal(response.Body.Bytes(), &users))
|
||||
require.Equal(t, map[string]string{"alice": "alice", "bob": "bob"}, users)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/lestrrat-go/jwx/v4/jwt"
|
||||
)
|
||||
|
||||
const (
|
||||
KnownBrowserLifetime = 180 * 24 * time.Hour
|
||||
knownBrowserJWTType = "known-browser"
|
||||
)
|
||||
|
||||
type browserTokenService interface {
|
||||
generate(userID string, lifetime time.Duration) (string, error)
|
||||
verify(token, userID string) error
|
||||
}
|
||||
|
||||
type browserTokens struct {
|
||||
signer SessionTokenService
|
||||
appURL string
|
||||
}
|
||||
|
||||
// generate creates a notification-only browser marker using the shared session signer
|
||||
func (s *browserTokens) generate(userID string, lifetime time.Duration) (string, error) {
|
||||
if s.signer == nil {
|
||||
return "", errors.New("session token signer is not initialized")
|
||||
}
|
||||
if userID == "" || lifetime <= 0 {
|
||||
return "", errors.New("user ID and positive lifetime are required for a known-browser token")
|
||||
}
|
||||
|
||||
// Bind recognition to one user and instance, with a purpose that access-token validation rejects
|
||||
now := time.Now()
|
||||
token, err := jwt.NewBuilder().
|
||||
Subject(userID).
|
||||
Issuer(s.appURL).
|
||||
Audience([]string{s.appURL}).
|
||||
IssuedAt(now).
|
||||
Expiration(now.Add(lifetime)).
|
||||
Claim("type", knownBrowserJWTType).
|
||||
Build()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to build known-browser token: %w", err)
|
||||
}
|
||||
|
||||
// The shared signer owns the private key and pinned signing algorithm
|
||||
return s.signer.SignSessionToken(token)
|
||||
}
|
||||
|
||||
// verify validates recognition without granting authentication or reading browser state
|
||||
func (s *browserTokens) verify(token, userID string) error {
|
||||
if s.signer == nil {
|
||||
return errors.New("session token signer is not initialized")
|
||||
}
|
||||
if userID == "" {
|
||||
return errors.New("user ID is required for a known-browser token")
|
||||
}
|
||||
|
||||
_, err := s.signer.VerifySessionToken(token,
|
||||
jwt.WithIssuer(s.appURL),
|
||||
jwt.WithAudience(s.appURL),
|
||||
jwt.WithSubject(userID),
|
||||
jwt.WithRequiredClaim(jwt.IssuedAtKey),
|
||||
jwt.WithRequiredClaim(jwt.ExpirationKey),
|
||||
jwt.WithClaimValue("type", knownBrowserJWTType),
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to verify known-browser token: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *service) rememberBrowser(userID, token string) (bool, string, error) {
|
||||
// Only a valid token for the authenticated user can suppress the new-browser notification
|
||||
known := token != "" && s.browserTokens.verify(token, userID) == nil
|
||||
|
||||
// Renew recognition after every successful sign-in without persisting browser state
|
||||
renewed, err := s.browserTokens.generate(userID, KnownBrowserLifetime)
|
||||
if err != nil {
|
||||
return known, token, err
|
||||
}
|
||||
return known, renewed, nil
|
||||
}
|
||||
@@ -0,0 +1,89 @@
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/lestrrat-go/jwx/v4/jwk"
|
||||
"github.com/lestrrat-go/jwx/v4/jwt"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
jwkutils "github.com/pocket-id/pocket-id/backend/internal/utils/jwk"
|
||||
)
|
||||
|
||||
type browserTokenTestSigner struct{ key jwk.Key }
|
||||
|
||||
func (s browserTokenTestSigner) SignSessionToken(token jwt.Token) (string, error) {
|
||||
signed, err := jwt.Sign(token, jwt.WithKey(jwkutils.SessionKeyAlg(), s.key))
|
||||
return string(signed), err
|
||||
}
|
||||
|
||||
func (s browserTokenTestSigner) VerifySessionToken(token string, options ...jwt.ValidateOption) (jwt.Token, error) {
|
||||
parseOptions := []jwt.ParseOption{jwt.WithKey(jwkutils.SessionKeyAlg(), s.key), jwt.WithValidate(true)}
|
||||
for _, option := range options {
|
||||
parseOptions = append(parseOptions, option)
|
||||
}
|
||||
return jwt.ParseString(token, parseOptions...)
|
||||
}
|
||||
|
||||
func newBrowserTokenService(t *testing.T) *browserTokens {
|
||||
t.Helper()
|
||||
key, err := jwkutils.GenerateSessionKey()
|
||||
require.NoError(t, err)
|
||||
return &browserTokens{signer: browserTokenTestSigner{key: key}, appURL: "https://test.example.com"}
|
||||
}
|
||||
|
||||
func TestKnownBrowserToken(t *testing.T) {
|
||||
s := newBrowserTokenService(t)
|
||||
const lifetime = 180 * 24 * time.Hour
|
||||
encoded, err := s.generate("user", lifetime)
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, s.verify(encoded, "user"))
|
||||
require.Error(t, s.verify(encoded, "other-user"))
|
||||
require.Error(t, s.verify(encoded, ""))
|
||||
|
||||
token, err := jwt.ParseString(encoded, jwt.WithKey(jwkutils.SessionKeyAlg(), s.signer.(browserTokenTestSigner).key))
|
||||
require.NoError(t, err)
|
||||
issued, ok := token.IssuedAt()
|
||||
require.True(t, ok)
|
||||
expires, ok := token.Expiration()
|
||||
require.True(t, ok)
|
||||
require.Equal(t, lifetime, expires.Sub(issued))
|
||||
}
|
||||
|
||||
func TestKnownBrowserTokenRejectsInvalidClaims(t *testing.T) {
|
||||
s := newBrowserTokenService(t)
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
claim string
|
||||
value any
|
||||
}{
|
||||
{"expired", jwt.ExpirationKey, time.Now().Add(-time.Second)},
|
||||
{"missing expiry", jwt.ExpirationKey, nil},
|
||||
{"missing issued at", jwt.IssuedAtKey, nil},
|
||||
{"future issued at", jwt.IssuedAtKey, time.Now().Add(time.Hour)},
|
||||
{"wrong issuer", jwt.IssuerKey, "https://other.example.com"},
|
||||
{"missing issuer", jwt.IssuerKey, nil},
|
||||
{"wrong audience", jwt.AudienceKey, []string{"other"}},
|
||||
{"missing audience", jwt.AudienceKey, nil},
|
||||
{"missing user", jwt.SubjectKey, nil},
|
||||
{"wrong purpose", "type", "access-token"},
|
||||
{"missing purpose", "type", nil},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Sign invalid claims with the real key so validation, rather than signature failure, rejects them
|
||||
token, err := jwt.NewBuilder().Subject("user").Issuer(s.appURL).
|
||||
Audience([]string{s.appURL}).IssuedAt(time.Now()).Expiration(time.Now().Add(time.Hour)).
|
||||
Claim("type", knownBrowserJWTType).Build()
|
||||
require.NoError(t, err)
|
||||
if tt.value == nil {
|
||||
require.NoError(t, token.Remove(tt.claim))
|
||||
} else {
|
||||
require.NoError(t, token.Set(tt.claim, tt.value))
|
||||
}
|
||||
encoded, err := jwt.Sign(token, jwt.WithKey(jwkutils.SessionKeyAlg(), s.signer.(browserTokenTestSigner).key))
|
||||
require.NoError(t, err)
|
||||
require.Error(t, s.verify(string(encoded), "user"))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
)
|
||||
|
||||
type AuditLog struct {
|
||||
model.Base
|
||||
|
||||
Event Event `sortable:"true" filterable:"true"`
|
||||
IpAddress *string `sortable:"true"`
|
||||
Country string `sortable:"true"`
|
||||
City string `sortable:"true"`
|
||||
UserAgent string `sortable:"true"`
|
||||
Username string `gorm:"-"`
|
||||
Data Data
|
||||
|
||||
UserID string `filterable:"true"`
|
||||
User model.User
|
||||
}
|
||||
|
||||
type Data map[string]string //nolint:recvcheck
|
||||
|
||||
type Event string //nolint:recvcheck
|
||||
|
||||
const (
|
||||
EventSignIn Event = "SIGN_IN"
|
||||
EventOneTimeAccessTokenSignIn Event = "TOKEN_SIGN_IN"
|
||||
EventRemoteSignIn Event = "REMOTE_SIGN_IN"
|
||||
EventAccountCreated Event = "ACCOUNT_CREATED"
|
||||
EventClientAuthorization Event = "CLIENT_AUTHORIZATION"
|
||||
EventNewClientAuthorization Event = "NEW_CLIENT_AUTHORIZATION"
|
||||
EventDeviceCodeAuthorization Event = "DEVICE_CODE_AUTHORIZATION"
|
||||
EventNewDeviceCodeAuthorization Event = "NEW_DEVICE_CODE_AUTHORIZATION"
|
||||
EventPasskeyAdded Event = "PASSKEY_ADDED"
|
||||
EventPasskeyRemoved Event = "PASSKEY_REMOVED"
|
||||
)
|
||||
|
||||
// Scan and Value methods for GORM to handle the custom type
|
||||
|
||||
func (e *Event) Scan(value any) error {
|
||||
*e = Event(value.(string))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e Event) Value() (driver.Value, error) {
|
||||
return string(e), nil
|
||||
}
|
||||
|
||||
func (d *Data) Scan(value any) error {
|
||||
return utils.UnmarshalJSONFromDatabase(d, value)
|
||||
}
|
||||
|
||||
func (d Data) Value() (driver.Value, error) {
|
||||
return json.Marshal(d)
|
||||
}
|
||||
@@ -1,44 +1,100 @@
|
||||
// Package auditlogs owns the background maintenance of the audit log table.
|
||||
// Package auditlogs owns audit records, their HTTP API, sign-in notifications, and retention cleanup
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
francishost "github.com/italypaleale/francis/host"
|
||||
"github.com/lestrrat-go/jwx/v4/jwt"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
|
||||
)
|
||||
|
||||
type NewLoginEmailSender interface {
|
||||
SendNewLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userFullName, userEmail, ipAddress, country, city, device, method string, dateTime time.Time) error
|
||||
}
|
||||
|
||||
type SessionTokenService interface {
|
||||
SignSessionToken(token jwt.Token) (string, error)
|
||||
VerifySessionToken(token string, options ...jwt.ValidateOption) (jwt.Token, error)
|
||||
}
|
||||
|
||||
type Dependencies struct {
|
||||
DB *gorm.DB
|
||||
Actors francishost.Host
|
||||
|
||||
EmailSender NewLoginEmailSender
|
||||
IPLocator iplocation.Resolver
|
||||
AppConfig appconfig.AppConfigResolver
|
||||
Signer SessionTokenService
|
||||
AppURL string
|
||||
|
||||
// RetentionDays is how long audit logs are kept before the cleanup job deletes them
|
||||
RetentionDays int
|
||||
|
||||
// CleanupDisabled skips registering the cleanup cron job, for example in tests
|
||||
// CleanupDisabled skips registering cleanup cron jobs, for example in tests
|
||||
CleanupDisabled bool
|
||||
}
|
||||
|
||||
type Module struct{}
|
||||
type Module struct {
|
||||
service *service
|
||||
handler *handler
|
||||
}
|
||||
|
||||
func New(deps Dependencies) (*Module, error) {
|
||||
// Register the cleanup job for audit logs past the retention window
|
||||
// Register audit retention cleanup before the actor host starts
|
||||
if !deps.CleanupDisabled {
|
||||
if deps.Actors == nil {
|
||||
return nil, errors.New("actor host is required for the audit log cleanup cron job")
|
||||
}
|
||||
|
||||
cleanupJob, err := newCleanupJob(deps.DB, deps.RetentionDays)
|
||||
jobs, err := newCleanupJobs(deps.DB, deps.RetentionDays)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
err = deps.Actors.RegisterBuiltInActor(cleanupJob)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("error registering audit log cleanup cron actor: %w", err)
|
||||
for _, cj := range jobs {
|
||||
if err := deps.Actors.RegisterBuiltInActor(cj); err != nil {
|
||||
return nil, fmt.Errorf("error registering audit log cleanup cron actor %q: %w", cj.ActorType(), err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return &Module{}, nil
|
||||
service := newService(deps.DB, deps.EmailSender, deps.IPLocator, deps.AppConfig, &browserTokens{signer: deps.Signer, appURL: deps.AppURL})
|
||||
return &Module{service: service, handler: newHandler(service)}, nil
|
||||
}
|
||||
|
||||
// RegisterRoutes mounts audit-log queries with the existing admin and current-user permissions
|
||||
func (m *Module) RegisterRoutes(group *gin.RouterGroup, adminAuth, userAuth gin.HandlerFunc) {
|
||||
group.GET("/audit-logs/all", adminAuth, httpserver.Handle(m.handler.listAllAuditLogsHandler))
|
||||
group.GET("/audit-logs", userAuth, httpserver.Handle(m.handler.listAuditLogsForUserHandler))
|
||||
group.GET("/audit-logs/filters/client-names", adminAuth, httpserver.Handle(m.handler.listClientNamesHandler))
|
||||
group.GET("/audit-logs/filters/users", adminAuth, httpserver.Handle(m.handler.listUserNamesWithIdsHandler))
|
||||
}
|
||||
|
||||
// Create records an event within the caller's transaction
|
||||
func (m *Module) Create(ctx context.Context, event Event, ipAddress, userAgent, userID string, data Data, tx *gorm.DB) (AuditLog, bool) {
|
||||
return m.service.Create(ctx, event, ipAddress, userAgent, userID, data, tx)
|
||||
}
|
||||
|
||||
// CreateSignIn prepares browser recognition and notification delivery within the caller's transaction
|
||||
func (m *Module) CreateSignIn(ctx context.Context, event Event, ipAddress, userAgent, userID, browserToken string, tx *gorm.DB, notificationMode appconfig.AppConfigValue) SignInResult {
|
||||
return m.service.CreateSignIn(ctx, event, ipAddress, userAgent, userID, browserToken, tx, notificationMode)
|
||||
}
|
||||
|
||||
// SendSignInNotification must be called only after the login commits successfully
|
||||
func (m *Module) SendSignInNotification(ctx context.Context, result SignInResult) {
|
||||
m.service.SendSignInNotification(ctx, result)
|
||||
}
|
||||
|
||||
// DeviceStringFromUserAgent describes a browser for audit records and device-login approval
|
||||
func (m *Module) DeviceStringFromUserAgent(userAgent string) string {
|
||||
return m.service.DeviceStringFromUserAgent(userAgent)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
userAgentParser "github.com/mileusna/useragent"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
)
|
||||
|
||||
type service struct {
|
||||
db *gorm.DB
|
||||
emailSender NewLoginEmailSender
|
||||
ipLocator iplocation.Resolver
|
||||
appConfigService appconfig.AppConfigResolver
|
||||
browserTokens browserTokenService
|
||||
}
|
||||
|
||||
func newService(db *gorm.DB, emailSender NewLoginEmailSender, ipLocator iplocation.Resolver, appConfigService appconfig.AppConfigResolver, browserTokens browserTokenService) *service {
|
||||
return &service{
|
||||
db: db,
|
||||
emailSender: emailSender,
|
||||
ipLocator: ipLocator,
|
||||
appConfigService: appConfigService,
|
||||
browserTokens: browserTokens,
|
||||
}
|
||||
}
|
||||
|
||||
// Create creates a new audit log entry in the database
|
||||
func (s *service) Create(ctx context.Context, event Event, ipAddress, userAgent, userID string, data Data, tx *gorm.DB) (AuditLog, bool) {
|
||||
country, city, err := s.ipLocator.GetLocationByIP(ctx, ipAddress)
|
||||
if err != nil {
|
||||
// Log the error but don't interrupt the operation
|
||||
slog.WarnContext(ctx, "Failed to get IP location", slog.String("ip", ipAddress), slog.Any("error", err))
|
||||
}
|
||||
|
||||
auditLog := AuditLog{
|
||||
Event: event,
|
||||
Country: country,
|
||||
City: city,
|
||||
UserAgent: userAgent,
|
||||
UserID: userID,
|
||||
Data: data,
|
||||
}
|
||||
|
||||
if ipAddress != "" {
|
||||
// Only set ipAddress if not empty, because on Postgres we use INET columns that don't allow non-null empty values
|
||||
auditLog.IpAddress = &ipAddress
|
||||
}
|
||||
|
||||
// Save the audit log in the database
|
||||
err = tx.
|
||||
WithContext(ctx).
|
||||
Create(&auditLog).
|
||||
Error
|
||||
if err != nil {
|
||||
slog.ErrorContext(ctx, "Failed to create audit log", "error", err)
|
||||
return AuditLog{}, false
|
||||
}
|
||||
|
||||
return auditLog, true
|
||||
}
|
||||
|
||||
// ListAuditLogsForUser retrieves all audit logs for a given user ID
|
||||
func (s *service) ListAuditLogsForUser(ctx context.Context, userID string, listRequestOptions utils.ListRequestOptions) ([]AuditLog, utils.PaginationResponse, error) {
|
||||
var logs []AuditLog
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Model(&AuditLog{}).
|
||||
Where("user_id = ?", userID)
|
||||
|
||||
pagination, err := utils.PaginateFilterAndSort(listRequestOptions, query, &logs)
|
||||
return logs, pagination, err
|
||||
}
|
||||
|
||||
func (s *service) DeviceStringFromUserAgent(userAgent string) string {
|
||||
ua := userAgentParser.Parse(userAgent)
|
||||
return ua.Name + " on " + ua.OS + " " + ua.OSVersion
|
||||
}
|
||||
|
||||
func (s *service) ListAllAuditLogs(ctx context.Context, listRequestOptions utils.ListRequestOptions) ([]AuditLog, utils.PaginationResponse, error) {
|
||||
var logs []AuditLog
|
||||
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Preload("User").
|
||||
Model(&AuditLog{})
|
||||
|
||||
if clientName, ok := listRequestOptions.Filters["clientName"]; ok {
|
||||
dialect := s.db.Name()
|
||||
switch dialect {
|
||||
case "sqlite":
|
||||
query = query.Where("json_extract(data, '$.clientName') IN ?", clientName)
|
||||
case "postgres":
|
||||
query = query.Where("data->>'clientName' IN ?", clientName)
|
||||
default:
|
||||
return nil, utils.PaginationResponse{}, fmt.Errorf("unsupported database dialect: %s", dialect)
|
||||
}
|
||||
}
|
||||
|
||||
if locations, ok := listRequestOptions.Filters["location"]; ok {
|
||||
mapped := make([]string, 0, len(locations))
|
||||
for _, v := range locations {
|
||||
if s, ok := v.(string); ok {
|
||||
switch s {
|
||||
case "internal":
|
||||
mapped = append(mapped, "Internal Network")
|
||||
case "external":
|
||||
mapped = append(mapped, "External Network")
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(mapped) > 0 {
|
||||
query = query.Where("country IN ?", mapped)
|
||||
}
|
||||
}
|
||||
|
||||
pagination, err := utils.PaginateFilterAndSort(listRequestOptions, query, &logs)
|
||||
if err != nil {
|
||||
return nil, pagination, err
|
||||
}
|
||||
|
||||
return logs, pagination, nil
|
||||
}
|
||||
|
||||
func (s *service) ListUsernamesWithIds(ctx context.Context) (users map[string]string, err error) {
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Joins("User").
|
||||
Model(&AuditLog{}).
|
||||
Select(`DISTINCT "User".id, "User".username`).
|
||||
Where(`"User".username IS NOT NULL`)
|
||||
|
||||
type Result struct {
|
||||
ID string `gorm:"column:id"`
|
||||
Username string `gorm:"column:username"`
|
||||
}
|
||||
|
||||
var results []Result
|
||||
err = query.Find(&results).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query user IDs: %w", err)
|
||||
}
|
||||
|
||||
users = make(map[string]string, len(results))
|
||||
for _, result := range results {
|
||||
users[result.ID] = result.Username
|
||||
}
|
||||
|
||||
return users, nil
|
||||
}
|
||||
|
||||
func (s *service) ListClientNames(ctx context.Context) (clientNames []string, err error) {
|
||||
dialect := s.db.Name()
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Model(&AuditLog{})
|
||||
|
||||
switch dialect {
|
||||
case "sqlite":
|
||||
query = query.
|
||||
Select("DISTINCT json_extract(data, '$.clientName') AS client_name").
|
||||
Where("json_extract(data, '$.clientName') IS NOT NULL")
|
||||
case "postgres":
|
||||
query = query.
|
||||
Select("DISTINCT data->>'clientName' AS client_name").
|
||||
Where("data->>'clientName' IS NOT NULL")
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported database dialect: %s", dialect)
|
||||
}
|
||||
|
||||
type Result struct {
|
||||
ClientName string `gorm:"column:client_name"`
|
||||
}
|
||||
|
||||
var results []Result
|
||||
err = query.Find(&results).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query client IDs: %w", err)
|
||||
}
|
||||
|
||||
clientNames = make([]string, len(results))
|
||||
for i, result := range results {
|
||||
clientNames[i] = result.ClientName
|
||||
}
|
||||
|
||||
return clientNames, nil
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
)
|
||||
|
||||
// CreateSignIn records a successful login and applies the global notification policy
|
||||
// The caller must send the notification only after the login commits
|
||||
func (s *service) CreateSignIn(ctx context.Context, event Event, ipAddress, userAgent, userID, browserToken string, tx *gorm.DB, notificationMode appconfig.AppConfigValue) SignInResult {
|
||||
entry, created := s.Create(ctx, event, ipAddress, userAgent, userID, Data{}, tx)
|
||||
result := SignInResult{AuditLog: entry, Created: created}
|
||||
if !created {
|
||||
return result
|
||||
}
|
||||
|
||||
// Only browser recognition uses a persistent cookie; other modes rely on the selected notification policy
|
||||
switch notificationMode {
|
||||
case appconfig.LoginNotificationAlways:
|
||||
result.Notify = true
|
||||
return result
|
||||
case appconfig.LoginNotificationBrowserRecognition:
|
||||
known, token, err := s.rememberBrowser(userID, browserToken)
|
||||
if err != nil {
|
||||
slog.ErrorContext(ctx, "Failed to remember sign-in browser", slog.Any("error", err))
|
||||
}
|
||||
result.KnownBrowserToken = token
|
||||
if known {
|
||||
return result
|
||||
}
|
||||
case appconfig.LoginNotificationIPAndUserAgent:
|
||||
// Check sign-in history without reading or renewing the browser token
|
||||
default:
|
||||
return result
|
||||
}
|
||||
|
||||
// Only earlier successful sign-ins for this user can satisfy the fallback
|
||||
var count int64
|
||||
query := tx.WithContext(ctx).Model(&AuditLog{}).
|
||||
Where("user_id = ? AND user_agent = ? AND id <> ?", userID, userAgent, entry.ID).
|
||||
Where("event IN ?", []Event{EventSignIn, EventOneTimeAccessTokenSignIn, EventRemoteSignIn})
|
||||
if ipAddress == "" {
|
||||
query = query.Where("ip_address IS NULL")
|
||||
} else {
|
||||
query = query.Where("ip_address = ?", ipAddress)
|
||||
}
|
||||
if err := query.Count(&count).Error; err != nil {
|
||||
slog.ErrorContext(ctx, "Failed to check sign-in history", slog.Any("error", err))
|
||||
return result
|
||||
}
|
||||
result.Notify = count == 0
|
||||
return result
|
||||
}
|
||||
|
||||
// SendSignInNotification sends only after the caller has completed the login successfully
|
||||
func (s *service) SendSignInNotification(ctx context.Context, result SignInResult) {
|
||||
if !result.Created || !result.Notify {
|
||||
return
|
||||
}
|
||||
entry := result.AuditLog
|
||||
ipAddress := ""
|
||||
if entry.IpAddress != nil {
|
||||
ipAddress = *entry.IpAddress
|
||||
}
|
||||
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)
|
||||
|
||||
// This runs after the request has completed, so we resolve the current config rather than threading the request's snapshot into the goroutine
|
||||
dbConfig, innerErr := s.appConfigService.GetConfig(innerCtx)
|
||||
if innerErr != nil {
|
||||
slog.ErrorContext(innerCtx, "Failed to load app configuration to send notification email", slog.Any("error", innerErr))
|
||||
return
|
||||
}
|
||||
|
||||
// Note we don't use the transaction here because this is running in background
|
||||
var user model.User
|
||||
innerErr = s.db.
|
||||
WithContext(innerCtx).
|
||||
Where("id = ?", entry.UserID).
|
||||
First(&user).
|
||||
Error
|
||||
if innerErr != nil {
|
||||
slog.ErrorContext(innerCtx, "Failed to load user from database to send notification email", slog.Any("error", innerErr))
|
||||
return
|
||||
}
|
||||
|
||||
if user.Email == nil {
|
||||
return
|
||||
}
|
||||
|
||||
innerErr = s.emailSender.SendNewLogin(
|
||||
innerCtx,
|
||||
dbConfig,
|
||||
user.FullName(),
|
||||
*user.Email,
|
||||
ipAddress,
|
||||
entry.Country,
|
||||
entry.City,
|
||||
s.DeviceStringFromUserAgent(entry.UserAgent),
|
||||
signInMethod(entry.Event),
|
||||
entry.CreatedAt.UTC(),
|
||||
)
|
||||
if innerErr != nil {
|
||||
slog.ErrorContext(innerCtx, "Failed to send notification email", slog.Any("error", innerErr), slog.String("address", *user.Email))
|
||||
return
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// signInMethod describes the successful login rather than the credential used to approve another device
|
||||
func signInMethod(event Event) string {
|
||||
switch event { //nolint:exhaustive // Other audit events are not sign-ins
|
||||
case EventSignIn:
|
||||
return "Passkey"
|
||||
case EventOneTimeAccessTokenSignIn:
|
||||
return "One-time code"
|
||||
case EventRemoteSignIn:
|
||||
return "Another device (QR code)"
|
||||
default:
|
||||
return "Unknown"
|
||||
}
|
||||
}
|
||||
|
||||
// SignInResult defers notification delivery until the login has committed
|
||||
type SignInResult struct {
|
||||
AuditLog AuditLog
|
||||
Created bool
|
||||
Notify bool
|
||||
KnownBrowserToken string
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
package auditlogs
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
)
|
||||
|
||||
type signInLocationResolver struct{}
|
||||
|
||||
func (signInLocationResolver) GetLocationByIP(context.Context, string) (string, string, error) {
|
||||
return "Switzerland", "Zurich", nil
|
||||
}
|
||||
|
||||
func TestSignInBrowserRecognition(t *testing.T) {
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
ip string
|
||||
agent string
|
||||
cookie string
|
||||
otherUser bool
|
||||
previousEvent Event
|
||||
wantNotify bool
|
||||
}{
|
||||
{name: "cookie recognizes changed IPv6 prefix and browser version", ip: "2001:db8:2::1", agent: "browser/2", cookie: "known"},
|
||||
{name: "missing cookie falls back to exact IP and agent", ip: "2001:db8:1::1", agent: "browser/1"},
|
||||
{name: "unknown cookie falls back to exact IP and agent", ip: "2001:db8:1::1", agent: "browser/1", cookie: "forged"},
|
||||
{name: "same IPv6 subnet is not an exact match", ip: "2001:db8:1::2", agent: "browser/1", wantNotify: true},
|
||||
{name: "same IP with changed agent is unknown", ip: "2001:db8:1::1", agent: "browser/2", wantNotify: true},
|
||||
{name: "forged cookie does not suppress notification", ip: "192.0.2.1", agent: "browser/1", cookie: "forged", wantNotify: true},
|
||||
{name: "cookie and history are scoped to user", ip: "2001:db8:1::1", agent: "browser/1", cookie: "known", otherUser: true, wantNotify: true},
|
||||
{name: "unrelated audit event does not establish familiarity", ip: "2001:db8:1::1", agent: "browser/1", previousEvent: EventClientAuthorization, wantNotify: true},
|
||||
{name: "login code history establishes familiarity", ip: "2001:db8:1::1", agent: "browser/1", previousEvent: EventOneTimeAccessTokenSignIn},
|
||||
{name: "QR login history establishes familiarity", ip: "2001:db8:1::1", agent: "browser/1", previousEvent: EventRemoteSignIn},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
s := newService(db, nil, signInLocationResolver{}, nil, browserTokenStub{})
|
||||
user := model.User{Base: model.Base{ID: "user"}, Username: "user"}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
historyUserID := user.ID
|
||||
if tt.otherUser {
|
||||
other := model.User{Base: model.Base{ID: "other"}, Username: "other"}
|
||||
require.NoError(t, db.Create(&other).Error)
|
||||
historyUserID = other.ID
|
||||
}
|
||||
event := tt.previousEvent
|
||||
if event == "" {
|
||||
event = EventSignIn
|
||||
}
|
||||
_, created := s.Create(t.Context(), event, "2001:db8:1::1", "browser/1", historyUserID, Data{}, db)
|
||||
require.True(t, created)
|
||||
s.browserTokens = browserTokenStub{userID: historyUserID}
|
||||
|
||||
result := s.CreateSignIn(t.Context(), EventSignIn, tt.ip, tt.agent, user.ID, tt.cookie, db, appconfig.LoginNotificationBrowserRecognition)
|
||||
require.True(t, result.Created)
|
||||
require.Equal(t, tt.wantNotify, result.Notify)
|
||||
|
||||
require.Equal(t, "renewed:"+user.ID, result.KnownBrowserToken)
|
||||
require.Equal(t, tt.ip, *result.AuditLog.IpAddress)
|
||||
require.Equal(t, tt.agent, result.AuditLog.UserAgent)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignInWithoutAddress(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
s := newService(db, nil, signInLocationResolver{}, nil, browserTokenStub{})
|
||||
user := model.User{Base: model.Base{ID: "user"}, Username: "user"}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
first := s.CreateSignIn(t.Context(), EventSignIn, "", "browser", user.ID, "", db, appconfig.LoginNotificationBrowserRecognition)
|
||||
require.True(t, first.Notify)
|
||||
require.Nil(t, first.AuditLog.IpAddress)
|
||||
second := s.CreateSignIn(t.Context(), EventSignIn, "", "browser", user.ID, "", db, appconfig.LoginNotificationBrowserRecognition)
|
||||
require.False(t, second.Notify)
|
||||
}
|
||||
|
||||
func TestSignInAuditRollsBackWithLogin(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
s := newService(db, nil, signInLocationResolver{}, nil, browserTokenStub{})
|
||||
user := model.User{Base: model.Base{ID: "user"}, Username: "user"}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
tx := db.Begin()
|
||||
require.NoError(t, tx.Error)
|
||||
result := s.CreateSignIn(t.Context(), EventSignIn, "192.0.2.1", "browser", user.ID, "", tx, appconfig.LoginNotificationBrowserRecognition)
|
||||
require.True(t, result.Created)
|
||||
require.True(t, result.Notify)
|
||||
require.NoError(t, tx.Rollback().Error)
|
||||
var count int64
|
||||
require.NoError(t, db.Model(&AuditLog{}).Count(&count).Error)
|
||||
require.Zero(t, count)
|
||||
}
|
||||
|
||||
type notificationConfig struct{ config *appconfig.AppConfigModel }
|
||||
|
||||
func (c notificationConfig) GetConfig(context.Context) (*appconfig.AppConfigModel, error) {
|
||||
return c.config, nil
|
||||
}
|
||||
|
||||
type loginNotification struct{ recipient, ip, method string }
|
||||
type loginNotificationSender struct{ sent chan loginNotification }
|
||||
|
||||
func (s loginNotificationSender) SendNewLogin(_ context.Context, _ *appconfig.AppConfigModel, _, email, ip, _, _, _, method string, _ time.Time) error {
|
||||
s.sent <- loginNotification{recipient: email, ip: ip, method: method}
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestSignInNotificationsForEveryMethod(t *testing.T) {
|
||||
for _, tt := range []struct {
|
||||
event Event
|
||||
method string
|
||||
}{
|
||||
{EventSignIn, "Passkey"},
|
||||
{EventOneTimeAccessTokenSignIn, "One-time code"},
|
||||
{EventRemoteSignIn, "Another device (QR code)"},
|
||||
} {
|
||||
t.Run(tt.method, func(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
email := "user@example.test"
|
||||
user := model.User{Base: model.Base{ID: "user"}, Username: "user", Email: &email}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
sender := loginNotificationSender{sent: make(chan loginNotification, 1)}
|
||||
config := &appconfig.AppConfigModel{EmailLoginNotificationMode: appconfig.LoginNotificationBrowserRecognition}
|
||||
service := newService(db, sender, signInLocationResolver{}, notificationConfig{config}, browserTokenStub{})
|
||||
|
||||
// Each successful sign-in method delivers a notification for an unfamiliar browser
|
||||
result := service.CreateSignIn(t.Context(), tt.event, "192.0.2.1", "browser", user.ID, "", db, appconfig.LoginNotificationBrowserRecognition)
|
||||
require.True(t, result.Created)
|
||||
require.True(t, result.Notify)
|
||||
service.SendSignInNotification(t.Context(), result)
|
||||
select {
|
||||
case sent := <-sender.sent:
|
||||
require.Equal(t, loginNotification{recipient: email, ip: "192.0.2.1", method: tt.method}, sent)
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("sign-in notification was not delivered")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// browserTokenStub exercises sign-in decisions while JWT validation is covered by the token service tests
|
||||
type browserTokenStub struct {
|
||||
userID string
|
||||
signErr error
|
||||
}
|
||||
|
||||
func (s browserTokenStub) generate(userID string, _ time.Duration) (string, error) {
|
||||
if s.signErr != nil {
|
||||
return "", s.signErr
|
||||
}
|
||||
return "renewed:" + userID, nil
|
||||
}
|
||||
|
||||
func (s browserTokenStub) verify(token, userID string) error {
|
||||
if token == "renewed:"+userID || (token == "known" && userID == s.userID) {
|
||||
return nil
|
||||
}
|
||||
return errors.New("invalid browser token")
|
||||
}
|
||||
|
||||
func TestSignInBrowserSigningFailurePreservesLogin(t *testing.T) {
|
||||
for _, token := range []string{"", "invalid-token", "known"} {
|
||||
t.Run(token, func(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
s := newService(db, nil, signInLocationResolver{}, nil, browserTokenStub{
|
||||
userID: "user", signErr: errors.New("signing failed"),
|
||||
})
|
||||
user := model.User{Base: model.Base{ID: "user"}, Username: "user"}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
tx := db.Begin()
|
||||
require.NoError(t, tx.Error)
|
||||
result := s.CreateSignIn(t.Context(), EventSignIn, "192.0.2.1", "browser", user.ID, token, tx, appconfig.LoginNotificationBrowserRecognition)
|
||||
require.NoError(t, tx.Commit().Error)
|
||||
require.True(t, result.Created)
|
||||
require.Equal(t, token, result.KnownBrowserToken)
|
||||
require.Equal(t, token != "known", result.Notify)
|
||||
require.NoError(t, db.First(&AuditLog{}, "id = ?", result.AuditLog.ID).Error)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSignInNotificationModes(t *testing.T) {
|
||||
for _, event := range []Event{EventSignIn, EventOneTimeAccessTokenSignIn, EventRemoteSignIn} {
|
||||
for _, tt := range []struct {
|
||||
mode appconfig.AppConfigValue
|
||||
knownHistory bool
|
||||
wantNotify bool
|
||||
}{
|
||||
{appconfig.LoginNotificationDisabled, false, false},
|
||||
{appconfig.LoginNotificationDisabled, true, false},
|
||||
{appconfig.LoginNotificationAlways, false, true},
|
||||
{appconfig.LoginNotificationAlways, true, true},
|
||||
{appconfig.LoginNotificationIPAndUserAgent, false, true},
|
||||
{appconfig.LoginNotificationIPAndUserAgent, true, false},
|
||||
{appconfig.LoginNotificationBrowserRecognition, false, false},
|
||||
{appconfig.LoginNotificationBrowserRecognition, true, false},
|
||||
} {
|
||||
t.Run(string(event)+"/"+string(tt.mode)+"/"+strconv.FormatBool(tt.knownHistory), func(t *testing.T) {
|
||||
db := testutils.NewDatabaseForTest(t)
|
||||
// A nil token service proves cookie-free modes never verify or issue browser tokens
|
||||
s := newService(db, nil, signInLocationResolver{}, nil, nil)
|
||||
if tt.mode == appconfig.LoginNotificationBrowserRecognition {
|
||||
s.browserTokens = browserTokenStub{userID: "user"}
|
||||
}
|
||||
user := model.User{Base: model.Base{ID: "user"}, Username: "user"}
|
||||
require.NoError(t, db.Create(&user).Error)
|
||||
if tt.knownHistory {
|
||||
_, created := s.Create(t.Context(), EventSignIn, "192.0.2.1", "browser", user.ID, Data{}, db)
|
||||
require.True(t, created)
|
||||
}
|
||||
result := s.CreateSignIn(t.Context(), event, "192.0.2.1", "browser", user.ID, "known", db, tt.mode)
|
||||
require.True(t, result.Created)
|
||||
require.Equal(t, tt.wantNotify, result.Notify)
|
||||
if tt.mode == appconfig.LoginNotificationBrowserRecognition {
|
||||
require.Equal(t, "renewed:user", result.KnownBrowserToken)
|
||||
} else {
|
||||
require.Empty(t, result.KnownBrowserToken)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -175,7 +175,7 @@ func registerRoutes(r *gin.Engine, db *gorm.DB, svc *services, rateLimitServices
|
||||
controller.NewAppConfigController(apiGroup, authMiddleware, svc.appConfigService, svc.emailModule)
|
||||
svc.ldapSyncModule.RegisterRoutes(apiGroup, authMiddleware.Add())
|
||||
controller.NewAppImagesController(apiGroup, authMiddleware, svc.appImagesService)
|
||||
controller.NewAuditLogController(apiGroup, svc.auditLogService, authMiddleware)
|
||||
svc.auditLogsModule.RegisterRoutes(apiGroup, authMiddleware.Add(), authMiddleware.WithAdminNotRequired().Add())
|
||||
controller.NewUserGroupController(apiGroup, authMiddleware, svc.appConfigService, svc.userGroupService)
|
||||
svc.apiModule.RegisterRoutes(apiGroup, authMiddleware.Add())
|
||||
controller.NewCustomClaimController(apiGroup, authMiddleware, svc.customClaimService)
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"net/http"
|
||||
|
||||
francishost "github.com/italypaleale/francis/host"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/api"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apikey"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
@@ -27,7 +29,6 @@ import (
|
||||
"github.com/pocket-id/pocket-id/backend/internal/storage"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/usersignup"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/webauthn"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type services struct {
|
||||
@@ -36,7 +37,6 @@ type services struct {
|
||||
emailModule *email.Module
|
||||
geoLiteModule *geolite.Module
|
||||
ipLocator iplocation.Resolver
|
||||
auditLogService *service.AuditLogService
|
||||
jwtService *service.JwtService
|
||||
userService *service.UserService
|
||||
customClaimService *service.CustomClaimService
|
||||
@@ -94,8 +94,17 @@ func initServices(
|
||||
return nil, fmt.Errorf("failed to create IP location resolver: %w", err)
|
||||
}
|
||||
|
||||
svc.auditLogService = service.NewAuditLogService(db, svc.emailModule, svc.ipLocator, svc.appConfigService)
|
||||
svc.jwtService, err = service.NewJwtService(ctx, db, instanceID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create JWT service: %w", err)
|
||||
}
|
||||
|
||||
svc.auditLogsModule, err = auditlogs.New(auditlogs.Dependencies{
|
||||
Signer: svc.jwtService,
|
||||
AppURL: common.EnvConfig.AppURL,
|
||||
EmailSender: svc.emailModule,
|
||||
IPLocator: svc.ipLocator,
|
||||
AppConfig: svc.appConfigService,
|
||||
DB: db,
|
||||
Actors: actors,
|
||||
RetentionDays: common.EnvConfig.AuditLogRetentionDays,
|
||||
@@ -106,18 +115,13 @@ func initServices(
|
||||
return nil, fmt.Errorf("failed to create audit logs module: %w", err)
|
||||
}
|
||||
|
||||
svc.jwtService, err = service.NewJwtService(ctx, db, instanceID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create JWT service: %w", err)
|
||||
}
|
||||
|
||||
svc.customClaimService = service.NewCustomClaimService(db)
|
||||
svc.webauthnModule, err = webauthn.New(webauthn.Dependencies{
|
||||
DB: db,
|
||||
Actors: actors,
|
||||
AppURL: common.EnvConfig.AppURL,
|
||||
Signer: svc.jwtService,
|
||||
AuditLog: svc.auditLogService,
|
||||
AuditLog: svc.auditLogsModule,
|
||||
AppConfig: svc.appConfigService,
|
||||
// Disable in test environment
|
||||
CleanupDisabled: common.EnvConfig.AppEnv.IsTest(),
|
||||
@@ -131,7 +135,7 @@ func initServices(
|
||||
Actors: actors,
|
||||
Signer: svc.jwtService,
|
||||
Reauth: svc.webauthnModule,
|
||||
AuditLog: svc.auditLogService,
|
||||
AuditLog: svc.auditLogsModule,
|
||||
IPLocator: svc.ipLocator,
|
||||
AppConfig: svc.appConfigService,
|
||||
})
|
||||
@@ -166,7 +170,7 @@ func initServices(
|
||||
Signer: svc.jwtService,
|
||||
CustomClaims: svc.customClaimService,
|
||||
Reauth: svc.webauthnModule,
|
||||
AuditLog: svc.auditLogService,
|
||||
AuditLog: svc.auditLogsModule,
|
||||
APIAccess: svc.apiModule,
|
||||
// Disable in test environment
|
||||
CleanupDisabled: common.EnvConfig.AppEnv.IsTest(),
|
||||
@@ -186,7 +190,7 @@ func initServices(
|
||||
}
|
||||
|
||||
svc.userGroupService = service.NewUserGroupService(db, svc.scimSyncModule, backchannelLogoutService)
|
||||
svc.userService = service.NewUserService(db, svc.jwtService, svc.auditLogService, svc.customClaimService, svc.appImagesService, svc.scimSyncModule, backchannelLogoutService, fileStorage)
|
||||
svc.userService = service.NewUserService(db, svc.jwtService, svc.customClaimService, svc.appImagesService, svc.scimSyncModule, backchannelLogoutService, fileStorage)
|
||||
|
||||
svc.ldapSyncModule, err = ldapsync.New(ldapsync.Dependencies{
|
||||
DB: db,
|
||||
@@ -221,7 +225,7 @@ func initServices(
|
||||
DB: db,
|
||||
Actors: actors,
|
||||
Signer: svc.jwtService,
|
||||
AuditLog: svc.auditLogService,
|
||||
AuditLog: svc.auditLogsModule,
|
||||
UserCreator: svc.userService,
|
||||
AppConfig: svc.appConfigService,
|
||||
ScimSync: svc.scimSyncModule,
|
||||
@@ -234,7 +238,7 @@ func initServices(
|
||||
DB: db,
|
||||
Actors: actors,
|
||||
Signer: svc.jwtService,
|
||||
AuditLog: svc.auditLogService,
|
||||
AuditLog: svc.auditLogsModule,
|
||||
UserProvider: svc.userService,
|
||||
EmailSender: svc.emailModule,
|
||||
AppConfig: svc.appConfigService,
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils/cookie"
|
||||
@@ -72,7 +73,8 @@ func (h *handler) exchangeRequest(c *gin.Context) error {
|
||||
requestID := c.Param("id")
|
||||
deviceToken, _ := c.Cookie(cookie.DeviceLoginTokenCookieName)
|
||||
sessionDuration := dbConfig.SessionDuration.AsDurationMinutes()
|
||||
user, accessToken, status, err := h.service.Exchange(c.Request.Context(), requestID, deviceToken, c.ClientIP(), c.Request.UserAgent(), sessionDuration)
|
||||
browserToken, _ := c.Cookie(cookie.KnownBrowserCookieName)
|
||||
user, tokens, status, err := h.service.Exchange(c.Request.Context(), requestID, deviceToken, c.ClientIP(), c.Request.UserAgent(), browserToken, sessionDuration, dbConfig.EmailLoginNotificationMode)
|
||||
if err != nil {
|
||||
if c.Request.Context().Err() != nil {
|
||||
// Context canceled = the client stopped the request
|
||||
@@ -88,7 +90,10 @@ func (h *handler) exchangeRequest(c *gin.Context) error {
|
||||
}
|
||||
|
||||
maxAge := int(sessionDuration.Seconds())
|
||||
cookie.AddAccessTokenCookie(c, maxAge, accessToken)
|
||||
cookie.AddAccessTokenCookie(c, maxAge, tokens.AccessToken)
|
||||
if tokens.KnownBrowserToken != "" {
|
||||
cookie.AddKnownBrowserCookie(c, tokens.KnownBrowserToken, int(auditlogs.KnownBrowserLifetime.Seconds()))
|
||||
}
|
||||
c.JSON(http.StatusOK, dto.UserDto(user))
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
@@ -24,7 +25,9 @@ type ReauthenticationTokenConsumer interface {
|
||||
}
|
||||
|
||||
type AuditLogger interface {
|
||||
Create(ctx context.Context, event model.AuditLogEvent, ipAddress, userAgent, userID string, data model.AuditLogData, tx *gorm.DB) (model.AuditLog, bool)
|
||||
CreateSignIn(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID, browserToken string, tx *gorm.DB, notificationMode appconfig.AppConfigValue) auditlogs.SignInResult
|
||||
SendSignInNotification(ctx context.Context, result auditlogs.SignInResult)
|
||||
Create(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID string, data auditlogs.Data, tx *gorm.DB) (auditlogs.AuditLog, bool)
|
||||
DeviceStringFromUserAgent(userAgent string) string
|
||||
}
|
||||
|
||||
|
||||
@@ -11,7 +11,9 @@ import (
|
||||
"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/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
@@ -153,9 +155,9 @@ func (s *Service) Decide(ctx context.Context, code, decision, userID, reauthenti
|
||||
return actorResultError(result.Code)
|
||||
}
|
||||
|
||||
func (s *Service) Exchange(ctx context.Context, requestID, deviceToken, ipAddress, userAgent string, sessionDuration time.Duration) (dto.UserDto, string, RequestStatus, error) {
|
||||
func (s *Service) Exchange(ctx context.Context, requestID, deviceToken, ipAddress, userAgent, browserToken string, sessionDuration time.Duration, notificationMode appconfig.AppConfigValue) (dto.UserDto, model.LoginTokens, RequestStatus, error) {
|
||||
if requestID == "" || deviceToken == "" || sessionDuration <= 0 {
|
||||
return dto.UserDto{}, "", "", apperror.DeviceLoginRequestInvalidOrExpired()
|
||||
return dto.UserDto{}, model.LoginTokens{}, "", apperror.DeviceLoginRequestInvalidOrExpired()
|
||||
}
|
||||
|
||||
deviceTokenHash := utils.CreateSha256Hash(deviceToken)
|
||||
@@ -168,12 +170,12 @@ func (s *Service) Exchange(ctx context.Context, requestID, deviceToken, ipAddres
|
||||
// Poll the actor's activation cache so the long-lived HTTP request does not repeatedly query the database
|
||||
result, err := s.peek(ctx, requestID, requestActorMethodPoll, requestActorPollInput{DeviceTokenHash: deviceTokenHash})
|
||||
if err != nil {
|
||||
return dto.UserDto{}, "", "", err
|
||||
return dto.UserDto{}, model.LoginTokens{}, "", err
|
||||
}
|
||||
|
||||
err = actorResultError(result.Code)
|
||||
if err != nil {
|
||||
return dto.UserDto{}, "", result.Status, err
|
||||
return dto.UserDto{}, model.LoginTokens{}, result.Status, err
|
||||
}
|
||||
|
||||
switch result.Status {
|
||||
@@ -181,7 +183,7 @@ func (s *Service) Exchange(ctx context.Context, requestID, deviceToken, ipAddres
|
||||
// Validate the approved user before consuming so lookup failures leave the request untouched
|
||||
user, userDTO, err := s.loadExchangeUser(ctx, result.UserID)
|
||||
if err != nil {
|
||||
return dto.UserDto{}, "", result.Status, err
|
||||
return dto.UserDto{}, model.LoginTokens{}, result.Status, err
|
||||
}
|
||||
|
||||
// Consume inside the actor so only one concurrent exchange can mint a token
|
||||
@@ -189,42 +191,43 @@ func (s *Service) Exchange(ctx context.Context, requestID, deviceToken, ipAddres
|
||||
DeviceTokenHash: deviceTokenHash,
|
||||
})
|
||||
if err != nil {
|
||||
return dto.UserDto{}, "", "", err
|
||||
return dto.UserDto{}, model.LoginTokens{}, "", err
|
||||
}
|
||||
|
||||
err = actorResultError(consume.Code)
|
||||
if err != nil {
|
||||
return dto.UserDto{}, "", consume.Status, err
|
||||
return dto.UserDto{}, model.LoginTokens{}, consume.Status, err
|
||||
}
|
||||
|
||||
// Mint the session with login-code semantics because the waiting device did not perform WebAuthn
|
||||
accessToken, err := s.signer.GenerateAccessToken(user, authenticationMethodOneTimePassword, sessionDuration)
|
||||
if err != nil {
|
||||
return dto.UserDto{}, "", consume.Status, err
|
||||
return dto.UserDto{}, model.LoginTokens{}, consume.Status, err
|
||||
}
|
||||
|
||||
// Record the successful remote sign-in after the request has been consumed
|
||||
_, created := s.auditLog.Create(ctx, model.AuditLogEventRemoteSignIn, ipAddress, userAgent, user.ID, model.AuditLogData{}, s.db)
|
||||
if !created {
|
||||
return dto.UserDto{}, "", consume.Status, errors.New("failed to create device login audit log")
|
||||
signIn := s.auditLog.CreateSignIn(ctx, auditlogs.EventRemoteSignIn, ipAddress, userAgent, user.ID, browserToken, s.db, notificationMode)
|
||||
if !signIn.Created {
|
||||
return dto.UserDto{}, model.LoginTokens{}, consume.Status, errors.New("failed to create device login audit log")
|
||||
}
|
||||
|
||||
return userDTO, accessToken, consume.Status, nil
|
||||
s.auditLog.SendSignInNotification(ctx, signIn)
|
||||
return userDTO, model.LoginTokens{AccessToken: accessToken, KnownBrowserToken: signIn.KnownBrowserToken}, consume.Status, nil
|
||||
case RequestStatusPending:
|
||||
// no-op
|
||||
case RequestStatusDenied:
|
||||
return dto.UserDto{}, "", result.Status, apperror.DeviceLoginDenied()
|
||||
return dto.UserDto{}, model.LoginTokens{}, result.Status, apperror.DeviceLoginDenied()
|
||||
default:
|
||||
return dto.UserDto{}, "", "", apperror.DeviceLoginRequestInvalidOrExpired()
|
||||
return dto.UserDto{}, model.LoginTokens{}, "", apperror.DeviceLoginRequestInvalidOrExpired()
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ticker.C:
|
||||
// no-op
|
||||
case <-timeout.C:
|
||||
return dto.UserDto{}, "", RequestStatusPending, nil
|
||||
return dto.UserDto{}, model.LoginTokens{}, RequestStatusPending, nil
|
||||
case <-ctx.Done():
|
||||
return dto.UserDto{}, "", "", ctx.Err()
|
||||
return dto.UserDto{}, model.LoginTokens{}, "", ctx.Err()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,7 +16,9 @@ import (
|
||||
"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/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"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"
|
||||
@@ -70,15 +72,16 @@ func (f *fakeTokenService) generatedToken() (string, string, time.Duration, int)
|
||||
}
|
||||
|
||||
type auditEntry struct {
|
||||
event model.AuditLogEvent
|
||||
event auditlogs.Event
|
||||
ipAddress string
|
||||
userAgent string
|
||||
userID string
|
||||
}
|
||||
|
||||
type fakeAuditLogger struct {
|
||||
mu sync.Mutex
|
||||
entries []auditEntry
|
||||
mu sync.Mutex
|
||||
entries []auditEntry
|
||||
notifications []auditlogs.SignInResult
|
||||
}
|
||||
|
||||
type fakeIPLocationResolver struct {
|
||||
@@ -91,11 +94,11 @@ func (f *fakeIPLocationResolver) GetLocationByIP(context.Context, string) (strin
|
||||
return f.country, f.city, f.err
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) Create(_ context.Context, event model.AuditLogEvent, ipAddress, userAgent, userID string, _ model.AuditLogData, _ *gorm.DB) (model.AuditLog, bool) {
|
||||
func (f *fakeAuditLogger) Create(_ context.Context, event auditlogs.Event, ipAddress, userAgent, userID string, _ auditlogs.Data, _ *gorm.DB) (auditlogs.AuditLog, bool) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.entries = append(f.entries, auditEntry{event: event, ipAddress: ipAddress, userAgent: userAgent, userID: userID})
|
||||
return model.AuditLog{}, true
|
||||
return auditlogs.AuditLog{}, true
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) DeviceStringFromUserAgent(userAgent string) string {
|
||||
@@ -154,11 +157,14 @@ func TestRequestLifecycle(t *testing.T) {
|
||||
err = fixture.service.Decide(t.Context(), strings.ToLower(request.Code), "approve", user.ID, "fresh-proof")
|
||||
require.NoError(t, err)
|
||||
|
||||
exchangedUser, accessToken, status, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "198.51.100.20", "target-agent", testSessionDuration)
|
||||
exchangedUser, accessToken, status, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "198.51.100.20", "target-agent", "", testSessionDuration, appconfig.LoginNotificationBrowserRecognition)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, RequestStatusApproved, status)
|
||||
require.Equal(t, user.ID, exchangedUser.ID)
|
||||
require.Equal(t, "device-login-access-token", accessToken)
|
||||
require.Equal(t, "device-login-access-token", accessToken.AccessToken)
|
||||
require.Len(t, fixture.auditLog.notifications, 1)
|
||||
require.True(t, fixture.auditLog.notifications[0].Notify)
|
||||
require.Equal(t, auditlogs.EventRemoteSignIn, fixture.auditLog.notifications[0].AuditLog.Event)
|
||||
|
||||
signedUserID, authenticationMethod, sessionDuration, generated := fixture.signer.generatedToken()
|
||||
require.Equal(t, user.ID, signedUserID)
|
||||
@@ -168,15 +174,15 @@ func TestRequestLifecycle(t *testing.T) {
|
||||
requireRequestActorStateDeleted(t, fixture.actors, request.ID)
|
||||
|
||||
entry := fixture.auditLog.lastEntry()
|
||||
require.Equal(t, model.AuditLogEventRemoteSignIn, entry.event)
|
||||
require.Equal(t, auditlogs.EventRemoteSignIn, entry.event)
|
||||
require.Equal(t, "198.51.100.20", entry.ipAddress)
|
||||
require.Equal(t, "target-agent", entry.userAgent)
|
||||
require.Equal(t, user.ID, entry.userID)
|
||||
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
assertInvalidRequestError(t, err)
|
||||
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, "wrong-token", "", "", testSessionDuration)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, "wrong-token", "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
assertInvalidRequestError(t, err)
|
||||
_, _, _, generated = fixture.signer.generatedToken()
|
||||
require.Equal(t, 1, generated)
|
||||
@@ -199,7 +205,7 @@ func TestPendingAndDeniedRequests(t *testing.T) {
|
||||
err = fixture.service.Decide(t.Context(), request.Code, "deny", "device-login-user", "")
|
||||
require.NoError(t, err)
|
||||
|
||||
user, accessToken, status, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
user, accessToken, status, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeDeviceLoginDenied))
|
||||
require.Equal(t, RequestStatusDenied, status)
|
||||
require.Empty(t, user.ID)
|
||||
@@ -219,7 +225,7 @@ func TestPendingExchangeObservesDecisionDuringLongPoll(t *testing.T) {
|
||||
}
|
||||
result := make(chan exchangeOutcome, 1)
|
||||
go func() {
|
||||
_, _, status, exchangeErr := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, status, exchangeErr := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
result <- exchangeOutcome{status: status, err: exchangeErr}
|
||||
}()
|
||||
|
||||
@@ -242,14 +248,14 @@ func TestRejectsInvalidAndExpiredRequestsWhileActorIsActive(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
unknownRequestID := strings.Repeat("a", 64)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), unknownRequestID, "device-token", "", "", testSessionDuration)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), unknownRequestID, "device-token", "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
assertInvalidRequestError(t, err)
|
||||
_, err = fixture.service.Inspect(t.Context(), unknownRequestID)
|
||||
assertInvalidRequestError(t, err)
|
||||
err = fixture.service.Decide(t.Context(), unknownRequestID, "deny", "device-login-user", "")
|
||||
assertInvalidRequestError(t, err)
|
||||
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, "wrong-token", "", "", testSessionDuration)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, "wrong-token", "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
assertInvalidRequestError(t, err)
|
||||
|
||||
state := getRequestActorState(t, fixture.actors, request.ID)
|
||||
@@ -260,7 +266,7 @@ func TestRejectsInvalidAndExpiredRequestsWhileActorIsActive(t *testing.T) {
|
||||
assertInvalidRequestError(t, err)
|
||||
err = fixture.service.Decide(t.Context(), request.Code, "deny", "device-login-user", "")
|
||||
assertInvalidRequestError(t, err)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
assertInvalidRequestError(t, err)
|
||||
}
|
||||
|
||||
@@ -279,7 +285,7 @@ func TestRejectsDisabledUserAtExchange(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, fixture.service.Decide(t.Context(), request.Code, "approve", user.ID, "fresh-proof"))
|
||||
|
||||
_, accessToken, _, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, accessToken, _, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeUserDisabled))
|
||||
require.Empty(t, accessToken)
|
||||
require.Equal(t, RequestStatusApproved, getRequestActorState(t, fixture.actors, request.ID).Status)
|
||||
@@ -300,14 +306,14 @@ func TestFailedTokenGenerationConsumesApprovedRequest(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, fixture.service.Decide(t.Context(), request.Code, "approve", user.ID, "fresh-proof"))
|
||||
|
||||
_, accessToken, status, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, accessToken, status, err := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
require.EqualError(t, err, "token generation failed")
|
||||
require.Empty(t, accessToken)
|
||||
require.Equal(t, RequestStatusApproved, status)
|
||||
requireRequestActorStateDeleted(t, fixture.actors, request.ID)
|
||||
require.Equal(t, 0, fixture.auditLog.entryCount())
|
||||
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, _, err = fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
assertInvalidRequestError(t, err)
|
||||
}
|
||||
|
||||
@@ -356,8 +362,8 @@ func TestConcurrentExchangeAllowsOnlyOneSuccess(t *testing.T) {
|
||||
waitGroup.Add(1)
|
||||
go func() {
|
||||
defer waitGroup.Done()
|
||||
_, token, _, exchangeErr := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
results <- exchangeResult{token: token, err: exchangeErr}
|
||||
_, token, _, exchangeErr := fixture.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
results <- exchangeResult{token: token.AccessToken, err: exchangeErr}
|
||||
}()
|
||||
}
|
||||
waitGroup.Wait()
|
||||
@@ -416,7 +422,7 @@ func TestRequestStateSurvivesActorHostRestart(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "persistent-agent", strings.TrimPrefix(info.Device, "Parsed "))
|
||||
require.NoError(t, secondModule.service.Decide(t.Context(), request.Code, "deny", "device-login-user", ""))
|
||||
_, _, status, err := secondModule.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, status, err := secondModule.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeDeviceLoginDenied))
|
||||
require.Equal(t, RequestStatusDenied, status)
|
||||
}
|
||||
@@ -435,14 +441,14 @@ func TestCompletedExchangeIsInvalidAfterActorHostRestart(t *testing.T) {
|
||||
request, deviceToken, err := firstModule.service.Create(t.Context(), "", "persistent-agent")
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, firstModule.service.Decide(t.Context(), request.Code, "approve", user.ID, "fresh-proof"))
|
||||
_, _, firstStatus, err := firstModule.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, firstStatus, err := firstModule.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, RequestStatusApproved, firstStatus)
|
||||
stopFirst()
|
||||
|
||||
secondModule, stopSecond := startPersistentDeviceLoginHost(t, db, deps)
|
||||
defer stopSecond()
|
||||
_, _, _, err = secondModule.service.Exchange(t.Context(), request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, _, err = secondModule.service.Exchange(t.Context(), request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
assertInvalidRequestError(t, err)
|
||||
|
||||
_, _, _, generated := deps.Signer.(*fakeTokenService).generatedToken()
|
||||
@@ -462,7 +468,7 @@ func TestPendingExchangeStopsWhenRequestIsCanceled(t *testing.T) {
|
||||
result := make(chan error, 1)
|
||||
go func() {
|
||||
close(started)
|
||||
_, _, _, exchangeErr := fixture.service.Exchange(ctx, request.ID, deviceToken, "", "", testSessionDuration)
|
||||
_, _, _, exchangeErr := fixture.service.Exchange(ctx, request.ID, deviceToken, "", "", "", testSessionDuration, appconfig.LoginNotificationDisabled)
|
||||
result <- exchangeErr
|
||||
}()
|
||||
|
||||
@@ -590,3 +596,15 @@ func freeLoopbackAddress(t *testing.T) string {
|
||||
require.NoError(t, listener.Close())
|
||||
return address
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) CreateSignIn(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID, browserToken string, tx *gorm.DB, mode appconfig.AppConfigValue) auditlogs.SignInResult {
|
||||
entry, created := f.Create(ctx, event, ipAddress, userAgent, userID, auditlogs.Data{}, tx)
|
||||
entry.Event = event
|
||||
return auditlogs.SignInResult{AuditLog: entry, Created: created, Notify: mode != appconfig.LoginNotificationDisabled, KnownBrowserToken: "recognized-browser"}
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) SendSignInNotification(_ context.Context, result auditlogs.SignInResult) {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
f.notifications = append(f.notifications, result)
|
||||
}
|
||||
|
||||
@@ -57,7 +57,7 @@ type AppConfigUpdateDto struct {
|
||||
WebauthnAuthenticatorAttachment string `json:"webauthnAuthenticatorAttachment" binding:"required,oneof=any platform cross-platform"`
|
||||
EmailOneTimeAccessAsAdminEnabled string `json:"emailOneTimeAccessAsAdminEnabled" binding:"required,boolean_string"`
|
||||
EmailOneTimeAccessAsUnauthenticatedEnabled string `json:"emailOneTimeAccessAsUnauthenticatedEnabled" binding:"required,boolean_string"`
|
||||
EmailLoginNotificationEnabled string `json:"emailLoginNotificationEnabled" binding:"required,boolean_string"`
|
||||
EmailLoginNotificationMode string `json:"emailLoginNotificationMode" binding:"required,oneof=disabled always ipAndUserAgent browserRecognition"`
|
||||
EmailApiKeyExpirationEnabled string `json:"emailApiKeyExpirationEnabled" binding:"required,boolean_string"`
|
||||
EmailVerificationEnabled string `json:"emailVerificationEnabled" binding:"required,boolean_string"`
|
||||
CIMDURLAllowlist string `json:"cimdUrlAllowlist" binding:"omitempty,cimd_url_allowlist"`
|
||||
|
||||
@@ -7,10 +7,11 @@ import (
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
datatype "github.com/pocket-id/pocket-id/backend/internal/model/types"
|
||||
)
|
||||
|
||||
type stringEnum string
|
||||
|
||||
type sourceStruct struct {
|
||||
AString string
|
||||
AStringPtr *string
|
||||
@@ -28,7 +29,7 @@ type sourceStruct struct {
|
||||
EmptyStringPtrToString *string
|
||||
NilStringPtrToString *string
|
||||
IntToInt64 int
|
||||
AuditLogEventToString model.AuditLogEvent
|
||||
StringEnumToString stringEnum
|
||||
}
|
||||
|
||||
type destStruct struct {
|
||||
@@ -48,7 +49,7 @@ type destStruct struct {
|
||||
EmptyStringPtrToString string
|
||||
NilStringPtrToString string
|
||||
IntToInt64 int64
|
||||
AuditLogEventToString string
|
||||
StringEnumToString string
|
||||
}
|
||||
|
||||
type embeddedStruct struct {
|
||||
@@ -83,7 +84,7 @@ func TestMapStruct(t *testing.T) {
|
||||
EmptyStringPtrToString: new(""),
|
||||
NilStringPtrToString: nil,
|
||||
IntToInt64: 99,
|
||||
AuditLogEventToString: model.AuditLogEventAccountCreated,
|
||||
StringEnumToString: stringEnum("example"),
|
||||
}
|
||||
var dst destStruct
|
||||
err := MapStruct(src, &dst)
|
||||
@@ -110,7 +111,7 @@ func TestMapStruct(t *testing.T) {
|
||||
assert.Empty(t, dst.EmptyStringPtrToString)
|
||||
assert.Empty(t, dst.NilStringPtrToString)
|
||||
assert.Equal(t, int64(99), dst.IntToInt64)
|
||||
assert.Equal(t, "ACCOUNT_CREATED", dst.AuditLogEventToString)
|
||||
assert.Equal(t, "example", dst.StringEnumToString)
|
||||
}
|
||||
|
||||
func TestMapStructList(t *testing.T) {
|
||||
|
||||
@@ -110,7 +110,7 @@ func (m *Module) SendOneTimeAccessEmail(ctx context.Context, dbConfig *appconfig
|
||||
})
|
||||
}
|
||||
|
||||
func (m *Module) SendNewLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userFullName, userEmail, ipAddress, country, city, device string, dateTime time.Time) error {
|
||||
func (m *Module) SendNewLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userFullName, userEmail, ipAddress, country, city, device, method string, dateTime time.Time) error {
|
||||
return send(ctx, m, dbConfig, address{
|
||||
name: userFullName,
|
||||
email: userEmail,
|
||||
@@ -119,6 +119,7 @@ func (m *Module) SendNewLogin(ctx context.Context, dbConfig *appconfig.AppConfig
|
||||
Country: country,
|
||||
City: city,
|
||||
Device: device,
|
||||
Method: method,
|
||||
DateTime: dateTime,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -83,9 +83,9 @@ func TestModuleSendsEveryEmailType(t *testing.T) {
|
||||
{
|
||||
name: "new login",
|
||||
subject: "New device login with Pocket ID Test",
|
||||
bodyContains: []string{"NEW SIGN-IN DETECTED", "Zurich, Switzerland", "192.0.2.10", "Firefox on Linux", "January 2, 2030 at 3:04 PM UTC"},
|
||||
bodyContains: []string{"NEW SIGN-IN DETECTED", "Zurich, Switzerland", "192.0.2.10", "Firefox on Linux", "Sign-in method", "One-time code", "January 2, 2030 at 3:04 PM UTC"},
|
||||
send: func(ctx context.Context, config *appconfig.AppConfigModel) error {
|
||||
return module.SendNewLogin(ctx, config, user.FullName(), userEmail, "192.0.2.10", "Switzerland", "Zurich", "Firefox on Linux", eventTime)
|
||||
return module.SendNewLogin(ctx, config, user.FullName(), userEmail, "192.0.2.10", "Switzerland", "Zurich", "Firefox on Linux", "One-time code", eventTime)
|
||||
},
|
||||
},
|
||||
{
|
||||
@@ -362,3 +362,17 @@ func readSMTPData(reader *bufio.Reader) (string, error) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewLoginRendersSignInMethodInBothTemplates(t *testing.T) {
|
||||
module, err := New(nil)
|
||||
require.NoError(t, err)
|
||||
text, html, err := renderBody(module, newLoginTemplate, &templateData[newLoginTemplateData]{
|
||||
AppName: "Pocket ID", AppURL: "https://id.example.test", LogoURL: "https://id.example.test/logo.png",
|
||||
Data: &newLoginTemplateData{IPAddress: "192.0.2.1", Device: "Firefox on Linux", Method: "One-time code", DateTime: time.Now()},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
for _, body := range []string{text, html} {
|
||||
assert.Contains(t, body, "Sign-in method")
|
||||
assert.Contains(t, body, "One-time code")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -47,6 +47,7 @@ type newLoginTemplateData struct {
|
||||
Country string
|
||||
City string
|
||||
Device string
|
||||
Method string
|
||||
DateTime time.Time
|
||||
}
|
||||
|
||||
|
||||
@@ -441,7 +441,6 @@ func newTestLdapService(t *testing.T, client ldapClient) (*Service, *gorm.DB) {
|
||||
userService := service.NewUserService(
|
||||
db,
|
||||
nil,
|
||||
nil,
|
||||
service.NewCustomClaimService(db),
|
||||
service.NewAppImagesService(map[string]string{}, fileStorage),
|
||||
nil,
|
||||
|
||||
@@ -39,7 +39,7 @@ func TestWithApiKeyAuthDisabled(t *testing.T) {
|
||||
jwtService, err := service.NewJwtService(t.Context(), db, instanceID)
|
||||
require.NoError(t, err)
|
||||
|
||||
userService := service.NewUserService(db, jwtService, nil, nil, nil, nil, nil, nil)
|
||||
userService := service.NewUserService(db, jwtService, nil, nil, nil, nil, nil)
|
||||
apiKeyModule, err := apikey.New(t.Context(), apikey.Dependencies{DB: db, CleanupDisabled: true})
|
||||
require.NoError(t, err)
|
||||
|
||||
|
||||
@@ -226,7 +226,7 @@ func TestValidationResponseUsesAppConfigTypeMessages(t *testing.T) {
|
||||
WebauthnAuthenticatorAttachment: "any",
|
||||
EmailOneTimeAccessAsAdminEnabled: "false",
|
||||
EmailOneTimeAccessAsUnauthenticatedEnabled: "false",
|
||||
EmailLoginNotificationEnabled: "false",
|
||||
EmailLoginNotificationMode: "disabled",
|
||||
EmailApiKeyExpirationEnabled: "false",
|
||||
EmailVerificationEnabled: "false",
|
||||
AutoCreateOIDCClientSecret: "true",
|
||||
|
||||
@@ -1,59 +0,0 @@
|
||||
package model
|
||||
|
||||
import (
|
||||
"database/sql/driver"
|
||||
"encoding/json"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
)
|
||||
|
||||
type AuditLog struct {
|
||||
Base
|
||||
|
||||
Event AuditLogEvent `sortable:"true" filterable:"true"`
|
||||
IpAddress *string `sortable:"true"`
|
||||
Country string `sortable:"true"`
|
||||
City string `sortable:"true"`
|
||||
UserAgent string `sortable:"true"`
|
||||
Username string `gorm:"-"`
|
||||
Data AuditLogData
|
||||
|
||||
UserID string `filterable:"true"`
|
||||
User User
|
||||
}
|
||||
|
||||
type AuditLogData map[string]string //nolint:recvcheck
|
||||
|
||||
type AuditLogEvent string //nolint:recvcheck
|
||||
|
||||
const (
|
||||
AuditLogEventSignIn AuditLogEvent = "SIGN_IN"
|
||||
AuditLogEventOneTimeAccessTokenSignIn AuditLogEvent = "TOKEN_SIGN_IN"
|
||||
AuditLogEventRemoteSignIn AuditLogEvent = "REMOTE_SIGN_IN"
|
||||
AuditLogEventAccountCreated AuditLogEvent = "ACCOUNT_CREATED"
|
||||
AuditLogEventClientAuthorization AuditLogEvent = "CLIENT_AUTHORIZATION"
|
||||
AuditLogEventNewClientAuthorization AuditLogEvent = "NEW_CLIENT_AUTHORIZATION"
|
||||
AuditLogEventDeviceCodeAuthorization AuditLogEvent = "DEVICE_CODE_AUTHORIZATION"
|
||||
AuditLogEventNewDeviceCodeAuthorization AuditLogEvent = "NEW_DEVICE_CODE_AUTHORIZATION"
|
||||
AuditLogEventPasskeyAdded AuditLogEvent = "PASSKEY_ADDED"
|
||||
AuditLogEventPasskeyRemoved AuditLogEvent = "PASSKEY_REMOVED"
|
||||
)
|
||||
|
||||
// Scan and Value methods for GORM to handle the custom type
|
||||
|
||||
func (e *AuditLogEvent) Scan(value any) error {
|
||||
*e = AuditLogEvent(value.(string))
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e AuditLogEvent) Value() (driver.Value, error) {
|
||||
return string(e), nil
|
||||
}
|
||||
|
||||
func (d *AuditLogData) Scan(value any) error {
|
||||
return utils.UnmarshalJSONFromDatabase(d, value)
|
||||
}
|
||||
|
||||
func (d AuditLogData) Value() (driver.Value, error) {
|
||||
return json.Marshal(d)
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
package model
|
||||
|
||||
// LoginTokens carries the separate authentication and browser recognition cookies
|
||||
type LoginTokens struct {
|
||||
AccessToken string
|
||||
KnownBrowserToken string
|
||||
}
|
||||
@@ -12,12 +12,14 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/ory/fosite"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"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"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
func newAuthorizationService(db *gorm.DB, interactionSessionService *interactionSessionService, claimsService *ClaimsService, reauth ReauthenticationTokenConsumer, auditLog AuditLogger, apiAccess APIAccessProvider) *authorizationService {
|
||||
@@ -277,12 +279,12 @@ func (s *authorizationService) authorizeAuthenticated(ctx context.Context, req a
|
||||
|
||||
grantResourceIndicator(req.requester, audience, grantedScopes)
|
||||
|
||||
authorizationEvent := model.AuditLogEventClientAuthorization
|
||||
authorizationEvent := auditlogs.EventClientAuthorization
|
||||
if !hasAlreadyAuthorizedClient {
|
||||
authorizationEvent = model.AuditLogEventNewClientAuthorization
|
||||
authorizationEvent = auditlogs.EventNewClientAuthorization
|
||||
}
|
||||
if s.auditLog != nil {
|
||||
s.auditLog.Create(ctx, authorizationEvent, req.meta.IPAddress, req.meta.UserAgent, req.userID, model.AuditLogData{"clientName": req.client.Name}, dbFromContext(ctx, s.db))
|
||||
s.auditLog.Create(ctx, authorizationEvent, req.meta.IPAddress, req.meta.UserAgent, req.userID, auditlogs.Data{"clientName": req.client.Name}, dbFromContext(ctx, s.db))
|
||||
}
|
||||
|
||||
return authorizationResult{Session: session}, nil
|
||||
@@ -691,7 +693,7 @@ func (s *authorizationService) completeConsentStep(ctx context.Context, interact
|
||||
return err
|
||||
}
|
||||
if !hasAlreadyAuthorizedClient && s.auditLog != nil {
|
||||
s.auditLog.Create(ctx, model.AuditLogEventNewClientAuthorization, meta.IPAddress, meta.UserAgent, userID, model.AuditLogData{"clientName": interactionSession.Client.Name}, dbFromContext(ctx, s.db))
|
||||
s.auditLog.Create(ctx, auditlogs.EventNewClientAuthorization, meta.IPAddress, meta.UserAgent, userID, auditlogs.Data{"clientName": interactionSession.Client.Name}, dbFromContext(ctx, s.db))
|
||||
}
|
||||
interactionSession.ConsentRequired = false
|
||||
return nil
|
||||
|
||||
@@ -8,23 +8,24 @@ import (
|
||||
|
||||
"github.com/ory/fosite"
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"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"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type fakeAuditLogger struct {
|
||||
events []model.AuditLogEvent
|
||||
data []model.AuditLogData
|
||||
events []auditlogs.Event
|
||||
data []auditlogs.Data
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) Create(_ context.Context, event model.AuditLogEvent, _, _, _ string, data model.AuditLogData, _ *gorm.DB) (model.AuditLog, bool) {
|
||||
func (f *fakeAuditLogger) Create(_ context.Context, event auditlogs.Event, _, _, _ string, data auditlogs.Data, _ *gorm.DB) (auditlogs.AuditLog, bool) {
|
||||
f.events = append(f.events, event)
|
||||
f.data = append(f.data, data)
|
||||
return model.AuditLog{}, true
|
||||
return auditlogs.AuditLog{}, true
|
||||
}
|
||||
|
||||
func TestAuthorizationServiceAuthorizeLogsClientAuthorization(t *testing.T) {
|
||||
@@ -59,8 +60,8 @@ func TestAuthorizationServiceAuthorizeLogsClientAuthorization(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
require.False(t, authorization.RequiresInteraction)
|
||||
|
||||
require.Equal(t, []model.AuditLogEvent{model.AuditLogEventClientAuthorization}, auditLogger.events)
|
||||
require.Equal(t, model.AuditLogData{"clientName": "Test Client"}, auditLogger.data[0])
|
||||
require.Equal(t, []auditlogs.Event{auditlogs.EventClientAuthorization}, auditLogger.events)
|
||||
require.Equal(t, auditlogs.Data{"clientName": "Test Client"}, auditLogger.data[0])
|
||||
}
|
||||
|
||||
func TestAuthorizationServiceRejectsCustomScopeWithoutResource(t *testing.T) {
|
||||
@@ -167,8 +168,8 @@ func TestAuthorizationServiceConsentStepLogsNewClientAuthorization(t *testing.T)
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, response.RedirectURL)
|
||||
|
||||
require.Equal(t, []model.AuditLogEvent{model.AuditLogEventNewClientAuthorization}, auditLogger.events)
|
||||
require.Equal(t, model.AuditLogData{"clientName": "Test Client"}, auditLogger.data[0])
|
||||
require.Equal(t, []auditlogs.Event{auditlogs.EventNewClientAuthorization}, auditLogger.events)
|
||||
require.Equal(t, auditlogs.Data{"clientName": "Test Client"}, auditLogger.data[0])
|
||||
}
|
||||
|
||||
func TestAuthorizationServiceConsentMergesAudienceQualifiedScopeKeys(t *testing.T) {
|
||||
@@ -234,7 +235,7 @@ func TestAuthorizationServiceConsentMergesAudienceQualifiedScopeKeys(t *testing.
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, authorization.RequiresInteraction)
|
||||
require.Equal(t, []model.AuditLogEvent{model.AuditLogEventClientAuthorization}, auditLogger.events)
|
||||
require.Equal(t, []auditlogs.Event{auditlogs.EventClientAuthorization}, auditLogger.events)
|
||||
}
|
||||
|
||||
// TestAuthorizationServiceRequiresConsentForScopelessAPIAccess guards that a token audienced to a custom API always needs
|
||||
@@ -1169,7 +1170,7 @@ func TestAuthorizationServiceSkipConsentGrantsWithoutInteraction(t *testing.T) {
|
||||
require.NoError(t, db.Model(&model.UserAuthorizedOidcClient{}).Where("user_id = ? AND client_id = ?", userID, clientID).Count(&count).Error)
|
||||
require.Equal(t, int64(1), count)
|
||||
|
||||
require.Equal(t, []model.AuditLogEvent{model.AuditLogEventNewClientAuthorization}, auditLogger.events)
|
||||
require.Equal(t, []auditlogs.Event{auditlogs.EventNewClientAuthorization}, auditLogger.events)
|
||||
}
|
||||
|
||||
// A client with SkipConsent must still show the consent screen when the request explicitly asks for it with prompt=consent
|
||||
|
||||
@@ -9,11 +9,13 @@ import (
|
||||
|
||||
"github.com/ory/fosite"
|
||||
"github.com/ory/fosite/handler/rfc8628"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"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"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type deviceService struct {
|
||||
@@ -143,11 +145,11 @@ func (s *deviceService) acceptDeviceCode(ctx context.Context, userCode, userID,
|
||||
return err
|
||||
}
|
||||
|
||||
event := model.AuditLogEventDeviceCodeAuthorization
|
||||
event := auditlogs.EventDeviceCodeAuthorization
|
||||
if !hasAlreadyAuthorizedClient {
|
||||
event = model.AuditLogEventNewDeviceCodeAuthorization
|
||||
event = auditlogs.EventNewDeviceCodeAuthorization
|
||||
}
|
||||
s.auditLog.Create(ctx, event, meta.IPAddress, meta.UserAgent, userID, model.AuditLogData{"clientName": client.Name}, dbFromContext(ctx, s.db))
|
||||
s.auditLog.Create(ctx, event, meta.IPAddress, meta.UserAgent, userID, auditlogs.Data{"clientName": client.Name}, dbFromContext(ctx, s.db))
|
||||
|
||||
deviceCodeSignature, err := s.store.AcceptDeviceCodeSessionByUserCodeSignature(ctx, userCodeSignature, request)
|
||||
if err != nil {
|
||||
|
||||
@@ -10,9 +10,11 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
francishost "github.com/italypaleale/francis/host"
|
||||
"github.com/lestrrat-go/jwx/v4/jwa"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
@@ -37,7 +39,7 @@ type ReauthenticationTokenConsumer interface {
|
||||
}
|
||||
|
||||
type AuditLogger interface {
|
||||
Create(ctx context.Context, event model.AuditLogEvent, ipAddress, userAgent, userID string, data model.AuditLogData, tx *gorm.DB) (model.AuditLog, bool)
|
||||
Create(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID string, data auditlogs.Data, tx *gorm.DB) (auditlogs.AuditLog, bool)
|
||||
}
|
||||
|
||||
type Dependencies struct {
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils/cookie"
|
||||
@@ -147,7 +148,8 @@ func (h *handler) exchangeToken(c *gin.Context) error {
|
||||
}
|
||||
|
||||
deviceToken, _ := c.Cookie(cookie.DeviceTokenCookieName)
|
||||
user, token, err := h.service.ExchangeToken(c.Request.Context(), cfg, loginCode, deviceToken, c.ClientIP(), c.Request.UserAgent())
|
||||
browserToken, _ := c.Cookie(cookie.KnownBrowserCookieName)
|
||||
user, tokens, err := h.service.ExchangeToken(c.Request.Context(), cfg, loginCode, deviceToken, c.ClientIP(), c.Request.UserAgent(), browserToken)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -159,7 +161,10 @@ func (h *handler) exchangeToken(c *gin.Context) error {
|
||||
}
|
||||
|
||||
maxAge := int(cfg.SessionDuration.AsDurationMinutes().Seconds())
|
||||
cookie.AddAccessTokenCookie(c, maxAge, token)
|
||||
cookie.AddAccessTokenCookie(c, maxAge, tokens.AccessToken)
|
||||
if tokens.KnownBrowserToken != "" {
|
||||
cookie.AddKnownBrowserCookie(c, tokens.KnownBrowserToken, int(auditlogs.KnownBrowserLifetime.Seconds()))
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, userDto)
|
||||
return nil
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
)
|
||||
@@ -24,7 +25,9 @@ type TokenService interface {
|
||||
}
|
||||
|
||||
type AuditLogger interface {
|
||||
Create(ctx context.Context, event model.AuditLogEvent, ipAddress, userAgent, userID string, data model.AuditLogData, tx *gorm.DB) (model.AuditLog, bool)
|
||||
CreateSignIn(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID, browserToken string, tx *gorm.DB, notificationMode appconfig.AppConfigValue) auditlogs.SignInResult
|
||||
SendSignInNotification(ctx context.Context, result auditlogs.SignInResult)
|
||||
Create(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID string, data auditlogs.Data, tx *gorm.DB) (auditlogs.AuditLog, bool)
|
||||
}
|
||||
|
||||
type UserProvider interface {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"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"
|
||||
@@ -152,7 +153,7 @@ func (s *Service) CreateToken(ctx context.Context, userID string, ttl time.Durat
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func (s *Service) ExchangeToken(ctx context.Context, dbConfig *appconfig.AppConfigModel, token, deviceToken, ipAddress, userAgent string) (model.User, string, error) {
|
||||
func (s *Service) ExchangeToken(ctx context.Context, dbConfig *appconfig.AppConfigModel, token, deviceToken, ipAddress, userAgent, browserToken string) (model.User, model.LoginTokens, error) {
|
||||
token = utils.NormalizeUnambiguousString(token)
|
||||
|
||||
// Consume the token by invoking its actor: this atomically validates it and, if valid, deletes it.
|
||||
@@ -161,38 +162,38 @@ func (s *Service) ExchangeToken(ctx context.Context, dbConfig *appconfig.AppConf
|
||||
DeviceToken: deviceToken,
|
||||
})
|
||||
if err != nil {
|
||||
return model.User{}, "", fmt.Errorf("error invoking one-time access token actor: %w", err)
|
||||
return model.User{}, model.LoginTokens{}, 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)
|
||||
return model.User{}, model.LoginTokens{}, fmt.Errorf("error decoding one-time access token actor response: %w", err)
|
||||
}
|
||||
|
||||
switch consumeRes.Status {
|
||||
case tokenConsumeNotFound:
|
||||
return model.User{}, "", apperror.TokenInvalidOrExpired()
|
||||
return model.User{}, model.LoginTokens{}, apperror.TokenInvalidOrExpired()
|
||||
case tokenConsumeDeviceMismatch:
|
||||
return model.User{}, "", apperror.DeviceCodeInvalid()
|
||||
return model.User{}, model.LoginTokens{}, apperror.DeviceCodeInvalid()
|
||||
case tokenConsumeOK:
|
||||
// All good, continue below
|
||||
default:
|
||||
return model.User{}, "", fmt.Errorf("unexpected status from one-time access token actor: %s", consumeRes.Status)
|
||||
return model.User{}, model.LoginTokens{}, 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)
|
||||
user, accessToken, err := s.completeTokenExchange(ctx, dbConfig, consumeRes.State, ipAddress, userAgent, browserToken)
|
||||
if err != nil {
|
||||
s.restoreToken(ctx, token, consumeRes.State)
|
||||
return model.User{}, "", err
|
||||
return model.User{}, model.LoginTokens{}, 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) {
|
||||
func (s *Service) completeTokenExchange(ctx context.Context, dbConfig *appconfig.AppConfigModel, state TokenState, ipAddress, userAgent, browserToken string) (model.User, model.LoginTokens, error) {
|
||||
var user model.User
|
||||
err := s.db.
|
||||
WithContext(ctx).
|
||||
@@ -200,13 +201,13 @@ func (s *Service) completeTokenExchange(ctx context.Context, dbConfig *appconfig
|
||||
First(&user).
|
||||
Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return model.User{}, "", apperror.TokenInvalidOrExpired()
|
||||
return model.User{}, model.LoginTokens{}, apperror.TokenInvalidOrExpired()
|
||||
} else if err != nil {
|
||||
return model.User{}, "", err
|
||||
return model.User{}, model.LoginTokens{}, err
|
||||
}
|
||||
|
||||
if user.Disabled {
|
||||
return model.User{}, "", apperror.UserDisabled()
|
||||
return model.User{}, model.LoginTokens{}, apperror.UserDisabled()
|
||||
}
|
||||
|
||||
accessToken, err := s.signer.GenerateAccessToken(
|
||||
@@ -215,18 +216,14 @@ func (s *Service) completeTokenExchange(ctx context.Context, dbConfig *appconfig
|
||||
dbConfig.SessionDuration.AsDurationMinutes(),
|
||||
)
|
||||
if err != nil {
|
||||
return model.User{}, "", err
|
||||
return model.User{}, model.LoginTokens{}, err
|
||||
}
|
||||
|
||||
s.auditLog.Create(
|
||||
ctx, model.AuditLogEventOneTimeAccessTokenSignIn,
|
||||
ipAddress, userAgent,
|
||||
user.ID,
|
||||
model.AuditLogData{},
|
||||
s.db,
|
||||
)
|
||||
// Recognize the receiving browser after the login code has been consumed and the user validated
|
||||
signIn := s.auditLog.CreateSignIn(ctx, auditlogs.EventOneTimeAccessTokenSignIn, ipAddress, userAgent, user.ID, browserToken, s.db, dbConfig.EmailLoginNotificationMode)
|
||||
s.auditLog.SendSignInNotification(ctx, signIn)
|
||||
|
||||
return user, accessToken, nil
|
||||
return user, model.LoginTokens{AccessToken: accessToken, KnownBrowserToken: signIn.KnownBrowserToken}, nil
|
||||
}
|
||||
|
||||
// restoreToken restores a token that was consumed but whose exchange could not be completed.
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
testutils "github.com/pocket-id/pocket-id/backend/internal/utils/testing"
|
||||
)
|
||||
@@ -23,12 +24,13 @@ func (fakeSigner) GenerateAccessToken(_ model.User, _ string, _ time.Duration) (
|
||||
}
|
||||
|
||||
type fakeAuditLogger struct {
|
||||
events []model.AuditLogEvent
|
||||
events []auditlogs.Event
|
||||
notifications []auditlogs.SignInResult
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) Create(_ context.Context, event model.AuditLogEvent, _, _, _ string, _ model.AuditLogData, _ *gorm.DB) (model.AuditLog, bool) {
|
||||
func (f *fakeAuditLogger) Create(_ context.Context, event auditlogs.Event, _, _, _ string, _ auditlogs.Data, _ *gorm.DB) (auditlogs.AuditLog, bool) {
|
||||
f.events = append(f.events, event)
|
||||
return model.AuditLog{}, true
|
||||
return auditlogs.AuditLog{}, true
|
||||
}
|
||||
|
||||
type fakeUserProvider struct {
|
||||
@@ -103,10 +105,14 @@ func TestExchangeTokenSuccess(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
exchangedUser, accessToken, err := svc.ExchangeToken(t.Context(), dbConfig, token, "", "1.2.3.4", "test-agent")
|
||||
dbConfig.EmailLoginNotificationMode = appconfig.LoginNotificationBrowserRecognition
|
||||
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)
|
||||
require.NotEmpty(t, accessToken.AccessToken)
|
||||
require.Len(t, auditLog.notifications, 1)
|
||||
require.True(t, auditLog.notifications[0].Notify)
|
||||
require.Equal(t, auditlogs.EventOneTimeAccessTokenSignIn, auditLog.notifications[0].AuditLog.Event)
|
||||
|
||||
// The token must have been consumed
|
||||
var state TokenState
|
||||
@@ -114,7 +120,7 @@ func TestExchangeTokenSuccess(t *testing.T) {
|
||||
require.ErrorIs(t, err, actor.ErrStateNotFound)
|
||||
|
||||
// A sign-in audit log must have been created
|
||||
require.Equal(t, []model.AuditLogEvent{model.AuditLogEventOneTimeAccessTokenSignIn}, auditLog.events)
|
||||
require.Equal(t, []auditlogs.Event{auditlogs.EventOneTimeAccessTokenSignIn}, auditLog.events)
|
||||
}
|
||||
|
||||
func TestExchangeTokenAcceptsAmbiguousAliases(t *testing.T) {
|
||||
@@ -134,7 +140,7 @@ func TestExchangeTokenAcceptsAmbiguousAliases(t *testing.T) {
|
||||
}, &actor.SetStateOpts{TTL: time.Minute}))
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
exchangedUser, _, err := svc.ExchangeToken(t.Context(), dbConfig, "aIObc2", "", "", "")
|
||||
exchangedUser, _, err := svc.ExchangeToken(t.Context(), dbConfig, "aIObc2", "", "", "", "")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, user.ID, exchangedUser.ID)
|
||||
}
|
||||
@@ -144,7 +150,7 @@ func TestExchangeTokenInvalidToken(t *testing.T) {
|
||||
svc, _, _ := newServiceForTest(t, db)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
_, _, err := svc.ExchangeToken(t.Context(), dbConfig, "does-not-exist", "", "", "")
|
||||
_, _, err := svc.ExchangeToken(t.Context(), dbConfig, "does-not-exist", "", "", "", "")
|
||||
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeTokenInvalidOrExpired))
|
||||
}
|
||||
@@ -165,7 +171,7 @@ func TestExchangeTokenDeviceMismatch(t *testing.T) {
|
||||
require.NotNil(t, deviceToken)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
_, _, err = svc.ExchangeToken(t.Context(), dbConfig, token, "wrong-device-token", "", "")
|
||||
_, _, err = svc.ExchangeToken(t.Context(), dbConfig, token, "wrong-device-token", "", "", "")
|
||||
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeDeviceCodeInvalid))
|
||||
|
||||
@@ -192,7 +198,7 @@ func TestExchangeTokenRejectsDisabledUser(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
dbConfig := appconfig.NewTestConfig(nil)
|
||||
exchangedUser, accessToken, err := svc.ExchangeToken(t.Context(), dbConfig, token, "", "", "")
|
||||
exchangedUser, accessToken, err := svc.ExchangeToken(t.Context(), dbConfig, token, "", "", "", "")
|
||||
|
||||
require.True(t, apperror.IsCode(err, apperror.CodeUserDisabled))
|
||||
require.Empty(t, exchangedUser.ID)
|
||||
@@ -205,4 +211,15 @@ func TestExchangeTokenRejectsDisabledUser(t *testing.T) {
|
||||
require.Equal(t, user.ID, state.UserID)
|
||||
|
||||
require.Empty(t, auditLog.events)
|
||||
require.Empty(t, auditLog.notifications)
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) CreateSignIn(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID, browserToken string, tx *gorm.DB, mode appconfig.AppConfigValue) auditlogs.SignInResult {
|
||||
entry, created := f.Create(ctx, event, ipAddress, userAgent, userID, auditlogs.Data{}, tx)
|
||||
entry.Event = event
|
||||
return auditlogs.SignInResult{AuditLog: entry, Created: created, Notify: mode != appconfig.LoginNotificationDisabled, KnownBrowserToken: "recognized-browser"}
|
||||
}
|
||||
|
||||
func (f *fakeAuditLogger) SendSignInNotification(_ context.Context, result auditlogs.SignInResult) {
|
||||
f.notifications = append(f.notifications, result)
|
||||
}
|
||||
|
||||
@@ -1,274 +0,0 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
userAgentParser "github.com/mileusna/useragent"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/iplocation"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type NewLoginEmailSender interface {
|
||||
SendNewLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, userFullName, userEmail, ipAddress, country, city, device string, dateTime time.Time) error
|
||||
}
|
||||
|
||||
type AuditLogService struct {
|
||||
db *gorm.DB
|
||||
emailSender NewLoginEmailSender
|
||||
ipLocator iplocation.Resolver
|
||||
appConfigService *appconfig.AppConfigService
|
||||
}
|
||||
|
||||
func NewAuditLogService(db *gorm.DB, emailSender NewLoginEmailSender, ipLocator iplocation.Resolver, appConfigService *appconfig.AppConfigService) *AuditLogService {
|
||||
return &AuditLogService{
|
||||
db: db,
|
||||
emailSender: emailSender,
|
||||
ipLocator: ipLocator,
|
||||
appConfigService: appConfigService,
|
||||
}
|
||||
}
|
||||
|
||||
// Create creates a new audit log entry in the database
|
||||
func (s *AuditLogService) Create(ctx context.Context, event model.AuditLogEvent, ipAddress, userAgent, userID string, data model.AuditLogData, tx *gorm.DB) (model.AuditLog, bool) {
|
||||
country, city, err := s.ipLocator.GetLocationByIP(ctx, ipAddress)
|
||||
if err != nil {
|
||||
// Log the error but don't interrupt the operation
|
||||
slog.WarnContext(ctx, "Failed to get IP location", slog.String("ip", ipAddress), slog.Any("error", err))
|
||||
}
|
||||
|
||||
auditLog := model.AuditLog{
|
||||
Event: event,
|
||||
Country: country,
|
||||
City: city,
|
||||
UserAgent: userAgent,
|
||||
UserID: userID,
|
||||
Data: data,
|
||||
}
|
||||
|
||||
if ipAddress != "" {
|
||||
// Only set ipAddress if not empty, because on Postgres we use INET columns that don't allow non-null empty values
|
||||
auditLog.IpAddress = &ipAddress
|
||||
}
|
||||
|
||||
// Save the audit log in the database
|
||||
err = tx.
|
||||
WithContext(ctx).
|
||||
Create(&auditLog).
|
||||
Error
|
||||
if err != nil {
|
||||
slog.Error("Failed to create audit log", "error", err)
|
||||
return model.AuditLog{}, false
|
||||
}
|
||||
|
||||
return auditLog, true
|
||||
}
|
||||
|
||||
// CreateNewSignInWithEmail creates a new audit log entry in the database and sends an email if the device hasn't been used before
|
||||
// emailLoginNotificationEnabled gates whether the notification email is sent, so the caller decides using the config it already loaded
|
||||
func (s *AuditLogService) CreateNewSignInWithEmail(ctx context.Context, ipAddress, userAgent, userID string, tx *gorm.DB, emailLoginNotificationEnabled bool) model.AuditLog {
|
||||
createdAuditLog, ok := s.Create(ctx, model.AuditLogEventSignIn, ipAddress, userAgent, userID, model.AuditLogData{}, tx)
|
||||
if !ok {
|
||||
// At this point the transaction has been canceled already, and error has been logged
|
||||
return createdAuditLog
|
||||
}
|
||||
|
||||
// Count the number of times the user has logged in from the same device
|
||||
var count int64
|
||||
stmt := tx.
|
||||
WithContext(ctx).
|
||||
Model(&model.AuditLog{}).
|
||||
Where("user_id = ? AND user_agent = ?", userID, userAgent)
|
||||
if ipAddress == "" {
|
||||
// An empty IP address is stored as NULL in the database
|
||||
stmt = stmt.Where("ip_address IS NULL")
|
||||
} else {
|
||||
stmt = stmt.Where("ip_address = ?", ipAddress)
|
||||
}
|
||||
err := stmt.Count(&count).Error
|
||||
if err != nil {
|
||||
slog.ErrorContext(ctx, "Failed to count audit logs", slog.Any("error", err))
|
||||
return createdAuditLog
|
||||
}
|
||||
|
||||
// If the user hasn't logged in from the same device before and email notifications are enabled, send an email
|
||||
if emailLoginNotificationEnabled && count <= 1 {
|
||||
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)
|
||||
|
||||
// This runs after the request has completed, so we resolve the current config rather than threading the request's snapshot into the goroutine
|
||||
dbConfig, innerErr := s.appConfigService.GetConfig(innerCtx)
|
||||
if innerErr != nil {
|
||||
slog.ErrorContext(innerCtx, "Failed to load app configuration to send notification email", slog.Any("error", innerErr))
|
||||
return
|
||||
}
|
||||
|
||||
// Note we don't use the transaction here because this is running in background
|
||||
var user model.User
|
||||
innerErr = s.db.
|
||||
WithContext(innerCtx).
|
||||
Where("id = ?", userID).
|
||||
First(&user).
|
||||
Error
|
||||
if innerErr != nil {
|
||||
slog.ErrorContext(innerCtx, "Failed to load user from database to send notification email", slog.Any("error", innerErr))
|
||||
return
|
||||
}
|
||||
|
||||
if user.Email == nil {
|
||||
return
|
||||
}
|
||||
|
||||
innerErr = s.emailSender.SendNewLogin(
|
||||
innerCtx,
|
||||
dbConfig,
|
||||
user.FullName(),
|
||||
*user.Email,
|
||||
ipAddress,
|
||||
createdAuditLog.Country,
|
||||
createdAuditLog.City,
|
||||
s.DeviceStringFromUserAgent(userAgent),
|
||||
createdAuditLog.CreatedAt.UTC(),
|
||||
)
|
||||
if innerErr != nil {
|
||||
slog.ErrorContext(innerCtx, "Failed to send notification email", slog.Any("error", innerErr), slog.String("address", *user.Email))
|
||||
return
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
return createdAuditLog
|
||||
}
|
||||
|
||||
// ListAuditLogsForUser retrieves all audit logs for a given user ID
|
||||
func (s *AuditLogService) ListAuditLogsForUser(ctx context.Context, userID string, listRequestOptions utils.ListRequestOptions) ([]model.AuditLog, utils.PaginationResponse, error) {
|
||||
var logs []model.AuditLog
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Model(&model.AuditLog{}).
|
||||
Where("user_id = ?", userID)
|
||||
|
||||
pagination, err := utils.PaginateFilterAndSort(listRequestOptions, query, &logs)
|
||||
return logs, pagination, err
|
||||
}
|
||||
|
||||
func (s *AuditLogService) DeviceStringFromUserAgent(userAgent string) string {
|
||||
ua := userAgentParser.Parse(userAgent)
|
||||
return ua.Name + " on " + ua.OS + " " + ua.OSVersion
|
||||
}
|
||||
|
||||
func (s *AuditLogService) ListAllAuditLogs(ctx context.Context, listRequestOptions utils.ListRequestOptions) ([]model.AuditLog, utils.PaginationResponse, error) {
|
||||
var logs []model.AuditLog
|
||||
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Preload("User").
|
||||
Model(&model.AuditLog{})
|
||||
|
||||
if clientName, ok := listRequestOptions.Filters["clientName"]; ok {
|
||||
dialect := s.db.Name()
|
||||
switch dialect {
|
||||
case "sqlite":
|
||||
query = query.Where("json_extract(data, '$.clientName') IN ?", clientName)
|
||||
case "postgres":
|
||||
query = query.Where("data->>'clientName' IN ?", clientName)
|
||||
default:
|
||||
return nil, utils.PaginationResponse{}, fmt.Errorf("unsupported database dialect: %s", dialect)
|
||||
}
|
||||
}
|
||||
|
||||
if locations, ok := listRequestOptions.Filters["location"]; ok {
|
||||
mapped := make([]string, 0, len(locations))
|
||||
for _, v := range locations {
|
||||
if s, ok := v.(string); ok {
|
||||
switch s {
|
||||
case "internal":
|
||||
mapped = append(mapped, "Internal Network")
|
||||
case "external":
|
||||
mapped = append(mapped, "External Network")
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(mapped) > 0 {
|
||||
query = query.Where("country IN ?", mapped)
|
||||
}
|
||||
}
|
||||
|
||||
pagination, err := utils.PaginateFilterAndSort(listRequestOptions, query, &logs)
|
||||
if err != nil {
|
||||
return nil, pagination, err
|
||||
}
|
||||
|
||||
return logs, pagination, nil
|
||||
}
|
||||
|
||||
func (s *AuditLogService) ListUsernamesWithIds(ctx context.Context) (users map[string]string, err error) {
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Joins("User").
|
||||
Model(&model.AuditLog{}).
|
||||
Select(`DISTINCT "User".id, "User".username`).
|
||||
Where(`"User".username IS NOT NULL`)
|
||||
|
||||
type Result struct {
|
||||
ID string `gorm:"column:id"`
|
||||
Username string `gorm:"column:username"`
|
||||
}
|
||||
|
||||
var results []Result
|
||||
err = query.Find(&results).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query user IDs: %w", err)
|
||||
}
|
||||
|
||||
users = make(map[string]string, len(results))
|
||||
for _, result := range results {
|
||||
users[result.ID] = result.Username
|
||||
}
|
||||
|
||||
return users, nil
|
||||
}
|
||||
|
||||
func (s *AuditLogService) ListClientNames(ctx context.Context) (clientNames []string, err error) {
|
||||
dialect := s.db.Name()
|
||||
query := s.db.
|
||||
WithContext(ctx).
|
||||
Model(&model.AuditLog{})
|
||||
|
||||
switch dialect {
|
||||
case "sqlite":
|
||||
query = query.
|
||||
Select("DISTINCT json_extract(data, '$.clientName') AS client_name").
|
||||
Where("json_extract(data, '$.clientName') IS NOT NULL")
|
||||
case "postgres":
|
||||
query = query.
|
||||
Select("DISTINCT data->>'clientName' AS client_name").
|
||||
Where("data->>'clientName' IS NOT NULL")
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported database dialect: %s", dialect)
|
||||
}
|
||||
|
||||
type Result struct {
|
||||
ClientName string `gorm:"column:client_name"`
|
||||
}
|
||||
|
||||
var results []Result
|
||||
err = query.Find(&results).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query client IDs: %w", err)
|
||||
}
|
||||
|
||||
clientNames = make([]string, len(results))
|
||||
for i, result := range results {
|
||||
clientNames[i] = result.ClientName
|
||||
}
|
||||
|
||||
return clientNames, nil
|
||||
}
|
||||
@@ -356,13 +356,7 @@ func (s *JwtService) GenerateAccessToken(user model.User, authenticationMethod s
|
||||
return "", fmt.Errorf("failed to set '%s' claim in token: %w", common.AuthenticationMethodsClaim, err)
|
||||
}
|
||||
|
||||
// Session tokens are signed with the symmetric session key
|
||||
signed, err := jwt.Sign(token, jwt.WithKey(jwkutils.SessionKeyAlg(), s.sessionKey))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to sign token: %w", err)
|
||||
}
|
||||
|
||||
return string(signed), nil
|
||||
return s.SignSessionToken(token)
|
||||
}
|
||||
|
||||
// GenerateLogoutToken creates a logout token for OIDC Back-Channel Logout 1.0
|
||||
@@ -402,25 +396,49 @@ func (s *JwtService) GenerateLogoutToken(userID string, clientID string) (string
|
||||
return string(signed), nil
|
||||
}
|
||||
|
||||
func (s *JwtService) VerifyAccessToken(tokenString string) (jwt.Token, error) {
|
||||
// SignSessionToken signs feature-owned claims with the private symmetric session key
|
||||
func (s *JwtService) SignSessionToken(token jwt.Token) (string, error) {
|
||||
if s.sessionKey == nil {
|
||||
return "", errors.New("session key is not initialized")
|
||||
}
|
||||
|
||||
signed, err := jwt.Sign(token, jwt.WithKey(jwkutils.SessionKeyAlg(), s.sessionKey))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to sign token: %w", err)
|
||||
}
|
||||
return string(signed), nil
|
||||
}
|
||||
|
||||
// VerifySessionToken pins signature verification while the caller supplies its feature's validation rules
|
||||
func (s *JwtService) VerifySessionToken(tokenString string, options ...jwt.ValidateOption) (jwt.Token, error) {
|
||||
if s.sessionKey == nil {
|
||||
return nil, errors.New("session key is not initialized")
|
||||
}
|
||||
|
||||
token, err := jwt.ParseString(
|
||||
tokenString,
|
||||
parseOptions := []jwt.ParseOption{
|
||||
jwt.WithValidate(true),
|
||||
jwt.WithKey(jwkutils.SessionKeyAlg(), s.sessionKey),
|
||||
}
|
||||
for _, option := range options {
|
||||
parseOptions = append(parseOptions, option)
|
||||
}
|
||||
token, err := jwt.ParseString(tokenString, parseOptions...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse token: %w", err)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func (s *JwtService) VerifyAccessToken(tokenString string) (jwt.Token, error) {
|
||||
if s.sessionKey == nil {
|
||||
return nil, errors.New("session key is not initialized")
|
||||
}
|
||||
return s.VerifySessionToken(tokenString,
|
||||
jwt.WithAcceptableSkew(clockSkew),
|
||||
jwt.WithAudience(s.envConfig.AppURL),
|
||||
jwt.WithIssuer(s.envConfig.AppURL),
|
||||
jwt.WithValidator(TokenTypeValidator(AccessTokenJWTType)),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse token: %w", err)
|
||||
}
|
||||
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// GetPublicJWK returns the JSON Web Key (JWK) for the public key.
|
||||
|
||||
@@ -31,7 +31,6 @@ import (
|
||||
type UserService struct {
|
||||
db *gorm.DB
|
||||
jwtService *JwtService
|
||||
auditLogService *AuditLogService
|
||||
customClaimService *CustomClaimService
|
||||
appImagesService *AppImagesService
|
||||
scimSyncScheduler ScimSyncScheduler
|
||||
@@ -39,11 +38,10 @@ type UserService struct {
|
||||
fileStorage storage.FileStorage
|
||||
}
|
||||
|
||||
func NewUserService(db *gorm.DB, jwtService *JwtService, auditLogService *AuditLogService, customClaimService *CustomClaimService, appImagesService *AppImagesService, scimSyncScheduler ScimSyncScheduler, backchannelLogout *backchannellogout.Service, fileStorage storage.FileStorage) *UserService {
|
||||
func NewUserService(db *gorm.DB, jwtService *JwtService, customClaimService *CustomClaimService, appImagesService *AppImagesService, scimSyncScheduler ScimSyncScheduler, backchannelLogout *backchannellogout.Service, fileStorage storage.FileStorage) *UserService {
|
||||
return &UserService{
|
||||
db: db,
|
||||
jwtService: jwtService,
|
||||
auditLogService: auditLogService,
|
||||
customClaimService: customClaimService,
|
||||
appImagesService: appImagesService,
|
||||
scimSyncScheduler: scimSyncScheduler,
|
||||
|
||||
@@ -26,7 +26,6 @@ func newTestUserService(t *testing.T) (*UserService, *UserGroupService) {
|
||||
userService := NewUserService(
|
||||
db,
|
||||
nil,
|
||||
nil,
|
||||
NewCustomClaimService(db),
|
||||
NewAppImagesService(map[string]string{}, fileStorage),
|
||||
nil,
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
@@ -20,7 +21,7 @@ type TokenService interface {
|
||||
}
|
||||
|
||||
type AuditLogger interface {
|
||||
Create(ctx context.Context, event model.AuditLogEvent, ipAddress, userAgent, userID string, data model.AuditLogData, tx *gorm.DB) (model.AuditLog, bool)
|
||||
Create(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID string, data auditlogs.Data, tx *gorm.DB) (auditlogs.AuditLog, bool)
|
||||
}
|
||||
|
||||
type UserCreator interface {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"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"
|
||||
@@ -115,11 +116,11 @@ func (s *Service) createSignedUpUser(ctx context.Context, config *appconfig.AppC
|
||||
}
|
||||
|
||||
if tokenProvided {
|
||||
s.auditLog.Create(ctx, model.AuditLogEventAccountCreated, ipAddress, userAgent, user.ID, model.AuditLogData{
|
||||
s.auditLog.Create(ctx, auditlogs.EventAccountCreated, ipAddress, userAgent, user.ID, auditlogs.Data{
|
||||
"signupToken": token,
|
||||
}, tx)
|
||||
} else {
|
||||
s.auditLog.Create(ctx, model.AuditLogEventAccountCreated, ipAddress, userAgent, user.ID, model.AuditLogData{
|
||||
s.auditLog.Create(ctx, auditlogs.EventAccountCreated, ipAddress, userAgent, user.ID, auditlogs.Data{
|
||||
"method": "open_signup",
|
||||
}, tx)
|
||||
}
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"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"
|
||||
@@ -37,8 +38,8 @@ func (fakeSigner) GenerateAccessToken(_ model.User, _ string, _ time.Duration) (
|
||||
|
||||
type fakeAuditLogger struct{}
|
||||
|
||||
func (fakeAuditLogger) Create(_ context.Context, _ model.AuditLogEvent, _, _, _ string, _ model.AuditLogData, _ *gorm.DB) (model.AuditLog, bool) {
|
||||
return model.AuditLog{}, true
|
||||
func (fakeAuditLogger) Create(_ context.Context, _ auditlogs.Event, _, _, _ string, _ auditlogs.Data, _ *gorm.DB) (auditlogs.AuditLog, bool) {
|
||||
return auditlogs.AuditLog{}, true
|
||||
}
|
||||
|
||||
func newSignupServiceForTest(t *testing.T, db *gorm.DB, userCreator UserCreator) *Service {
|
||||
|
||||
@@ -32,3 +32,7 @@ func addCookie(c *gin.Context, name, value string, maxAge int, path string) {
|
||||
c.SetSameSite(http.SameSiteLaxMode)
|
||||
c.SetCookie(name, value, maxAge, path, "", true, true)
|
||||
}
|
||||
|
||||
func AddKnownBrowserCookie(c *gin.Context, token string, maxAge int) {
|
||||
addCookie(c, KnownBrowserCookieName, token, maxAge, "/")
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"github.com/pocket-id/pocket-id/backend/internal/common"
|
||||
)
|
||||
|
||||
var KnownBrowserCookieName = "__Host-known_browser"
|
||||
var AccessTokenCookieName = "__Host-access_token"
|
||||
var SessionIdCookieName = "__Host-session"
|
||||
var DeviceTokenCookieName = "__Secure-device_token" // #nosec G101 -- cookie name, not a credential
|
||||
@@ -14,6 +15,7 @@ var ReauthenticationTokenCookieName = "__Secure-reauthentication_token" // #nose
|
||||
|
||||
func init() {
|
||||
if strings.HasPrefix(common.EnvConfig.AppURL, "http://") {
|
||||
KnownBrowserCookieName = "known_browser"
|
||||
AccessTokenCookieName = "access_token"
|
||||
SessionIdCookieName = "session"
|
||||
DeviceTokenCookieName = "device_token"
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/dto"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/utils"
|
||||
@@ -106,7 +107,8 @@ func (h *handler) verifyLogin(c *gin.Context) error {
|
||||
return apperror.InvalidWebAuthnResponse(err)
|
||||
}
|
||||
|
||||
user, token, err := h.service.VerifyLogin(c.Request.Context(), dbConfig, sessionID, credentialAssertionData, c.ClientIP(), c.Request.UserAgent())
|
||||
browserToken, _ := c.Cookie(cookie.KnownBrowserCookieName)
|
||||
user, tokens, err := h.service.VerifyLogin(c.Request.Context(), dbConfig, sessionID, credentialAssertionData, c.ClientIP(), c.Request.UserAgent(), browserToken)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -117,7 +119,10 @@ func (h *handler) verifyLogin(c *gin.Context) error {
|
||||
}
|
||||
|
||||
maxAge := int(dbConfig.SessionDuration.AsDurationMinutes().Seconds())
|
||||
cookie.AddAccessTokenCookie(c, maxAge, token)
|
||||
cookie.AddAccessTokenCookie(c, maxAge, tokens.AccessToken)
|
||||
if tokens.KnownBrowserToken != "" {
|
||||
cookie.AddKnownBrowserCookie(c, tokens.KnownBrowserToken, int(auditlogs.KnownBrowserLifetime.Seconds()))
|
||||
}
|
||||
|
||||
c.JSON(http.StatusOK, userDto)
|
||||
return nil
|
||||
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/httpserver"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/model"
|
||||
)
|
||||
@@ -23,8 +24,9 @@ type TokenService interface {
|
||||
}
|
||||
|
||||
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, emailLoginNotificationEnabled bool) model.AuditLog
|
||||
CreateSignIn(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID, browserToken string, tx *gorm.DB, notificationMode appconfig.AppConfigValue) auditlogs.SignInResult
|
||||
SendSignInNotification(ctx context.Context, result auditlogs.SignInResult)
|
||||
Create(ctx context.Context, event auditlogs.Event, ipAddress, userAgent, userID string, data auditlogs.Data, tx *gorm.DB) (auditlogs.AuditLog, bool)
|
||||
}
|
||||
|
||||
type Dependencies struct {
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
|
||||
"github.com/pocket-id/pocket-id/backend/internal/appconfig"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/apperror"
|
||||
"github.com/pocket-id/pocket-id/backend/internal/auditlogs"
|
||||
"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"
|
||||
@@ -204,8 +205,8 @@ func (s *Service) VerifyRegistration(ctx context.Context, dbConfig *appconfig.Ap
|
||||
return model.WebauthnCredential{}, fmt.Errorf("failed to store WebAuthn credential: %w", err)
|
||||
}
|
||||
|
||||
auditLogData := model.AuditLogData{"credentialID": hex.EncodeToString(credential.ID), "passkeyName": passkeyName}
|
||||
s.auditLog.Create(ctx, model.AuditLogEventPasskeyAdded, ipAddress, r.UserAgent(), userID, auditLogData, tx)
|
||||
auditLogData := auditlogs.Data{"credentialID": hex.EncodeToString(credential.ID), "passkeyName": passkeyName}
|
||||
s.auditLog.Create(ctx, auditlogs.EventPasskeyAdded, ipAddress, r.UserAgent(), userID, auditLogData, tx)
|
||||
|
||||
err = tx.Commit().Error
|
||||
if err != nil {
|
||||
@@ -255,7 +256,7 @@ func (s *Service) BeginLogin(ctx context.Context, dbConfig *appconfig.AppConfigM
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) VerifyLogin(ctx context.Context, dbConfig *appconfig.AppConfigModel, 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, browserToken string) (model.User, model.LoginTokens, error) {
|
||||
tx := s.db.Begin()
|
||||
defer func() {
|
||||
tx.Rollback()
|
||||
@@ -268,10 +269,10 @@ func (s *Service) VerifyLogin(ctx context.Context, dbConfig *appconfig.AppConfig
|
||||
Clauses(clause.Returning{}).
|
||||
Delete(&storedSession, "id = ?", sessionID)
|
||||
if result.Error != nil {
|
||||
return model.User{}, "", fmt.Errorf("failed to load WebAuthn session: %w", result.Error)
|
||||
return model.User{}, model.LoginTokens{}, fmt.Errorf("failed to load WebAuthn session: %w", result.Error)
|
||||
}
|
||||
if result.RowsAffected == 0 {
|
||||
return model.User{}, "", apperror.InvalidWebAuthnSession()
|
||||
return model.User{}, model.LoginTokens{}, apperror.InvalidWebAuthnSession()
|
||||
}
|
||||
|
||||
session := gowebauthn.SessionData{
|
||||
@@ -301,29 +302,35 @@ func (s *Service) VerifyLogin(ctx context.Context, dbConfig *appconfig.AppConfig
|
||||
}, session, credentialAssertionData)
|
||||
|
||||
if err != nil {
|
||||
return model.User{}, "", classifyPasskeyError(err, apperror.WebAuthnAuthenticationFailed)
|
||||
return model.User{}, model.LoginTokens{}, classifyPasskeyError(err, apperror.WebAuthnAuthenticationFailed)
|
||||
}
|
||||
if user == nil {
|
||||
return model.User{}, "", apperror.WebAuthnAuthenticationFailed(errors.New("WebAuthn response did not resolve to a user"))
|
||||
return model.User{}, model.LoginTokens{}, apperror.WebAuthnAuthenticationFailed(errors.New("WebAuthn response did not resolve to a user"))
|
||||
}
|
||||
|
||||
if user.Disabled {
|
||||
return model.User{}, "", apperror.UserDisabled()
|
||||
return model.User{}, model.LoginTokens{}, apperror.UserDisabled()
|
||||
}
|
||||
|
||||
token, err := s.signer.GenerateAccessToken(*user, authenticationMethodPhishingResistant, dbConfig.SessionDuration.AsDurationMinutes())
|
||||
if err != nil {
|
||||
return model.User{}, "", err
|
||||
return model.User{}, model.LoginTokens{}, err
|
||||
}
|
||||
|
||||
s.auditLog.CreateNewSignInWithEmail(ctx, ipAddress, userAgent, user.ID, tx, dbConfig.EmailLoginNotificationEnabled.IsTrue())
|
||||
// Prepare browser recognition and the notification within the login transaction
|
||||
signIn := s.auditLog.CreateSignIn(ctx, auditlogs.EventSignIn, ipAddress, userAgent, user.ID, browserToken, tx, dbConfig.EmailLoginNotificationMode)
|
||||
if !signIn.Created {
|
||||
return model.User{}, model.LoginTokens{}, errors.New("failed to create sign-in audit log")
|
||||
}
|
||||
|
||||
err = tx.Commit().Error
|
||||
if err != nil {
|
||||
return model.User{}, "", err
|
||||
return model.User{}, model.LoginTokens{}, err
|
||||
}
|
||||
|
||||
return *user, token, nil
|
||||
// Deliver the notification only after the login transaction has committed
|
||||
s.auditLog.SendSignInNotification(ctx, signIn)
|
||||
return *user, model.LoginTokens{AccessToken: token, KnownBrowserToken: signIn.KnownBrowserToken}, nil
|
||||
}
|
||||
|
||||
func (s *Service) ListCredentials(ctx context.Context, userID string) ([]model.WebauthnCredential, error) {
|
||||
@@ -356,7 +363,7 @@ func (s *Service) DeleteCredential(ctx context.Context, userID string, credentia
|
||||
return apperror.NotFound("Passkey")
|
||||
}
|
||||
|
||||
auditLogData := model.AuditLogData{"credentialID": hex.EncodeToString(credential.CredentialID), "passkeyName": credential.Name}
|
||||
auditLogData := auditlogs.Data{"credentialID": hex.EncodeToString(credential.CredentialID), "passkeyName": credential.Name}
|
||||
if actorUserID != "" && actorUserID != userID {
|
||||
var actor model.User
|
||||
err := tx.
|
||||
@@ -372,7 +379,7 @@ func (s *Service) DeleteCredential(ctx context.Context, userID string, credentia
|
||||
auditLogData["actorUserID"] = actorUserID
|
||||
auditLogData["actorUsername"] = actor.Username
|
||||
}
|
||||
s.auditLog.Create(ctx, model.AuditLogEventPasskeyRemoved, ipAddress, userAgent, userID, auditLogData, tx)
|
||||
s.auditLog.Create(ctx, auditlogs.EventPasskeyRemoved, ipAddress, userAgent, userID, auditLogData, tx)
|
||||
|
||||
err := tx.Commit().Error
|
||||
if err != nil {
|
||||
|
||||
@@ -371,7 +371,7 @@ func TestCeremoniesRejectSessionThatDoesNotExist(t *testing.T) {
|
||||
t.Run("login rejects an unknown session", func(t *testing.T) {
|
||||
service := setupService(t)
|
||||
|
||||
_, token, err := service.VerifyLogin(t.Context(), &appconfig.AppConfigModel{}, "does-not-exist", nil, "127.0.0.1", "test-agent")
|
||||
_, token, err := service.VerifyLogin(t.Context(), &appconfig.AppConfigModel{}, "does-not-exist", nil, "127.0.0.1", "test-agent", "")
|
||||
|
||||
assert.Empty(t, token)
|
||||
require.Error(t, err)
|
||||
|
||||
File diff suppressed because one or more lines are too long
@@ -11,6 +11,9 @@ IP address
|
||||
Device
|
||||
{{.Data.Device}}
|
||||
|
||||
Sign-in method
|
||||
{{.Data.Method}}
|
||||
|
||||
Time
|
||||
{{.Data.DateTime.Format "January 2, 2006 at 3:04 PM MST"}}
|
||||
|
||||
|
||||
@@ -8,6 +8,7 @@ interface SignInData {
|
||||
location: string;
|
||||
ipAddress: string;
|
||||
device: string;
|
||||
method: string;
|
||||
dateTime: string;
|
||||
}
|
||||
|
||||
@@ -28,6 +29,7 @@ export const NewSignInEmail = ({ data, ...props }: NewSignInEmailProps) => (
|
||||
{ label: 'Approximate location', value: data.location },
|
||||
{ label: 'IP address', value: data.ipAddress },
|
||||
{ label: 'Device', value: data.device },
|
||||
{ label: 'Sign-in method', value: data.method },
|
||||
{ label: 'Time', value: data.dateTime }
|
||||
]}
|
||||
/>
|
||||
@@ -48,6 +50,7 @@ NewSignInEmail.TemplateProps = {
|
||||
'{{if and .Data.City .Data.Country}}{{.Data.City}}, {{.Data.Country}}{{else if .Data.Country}}{{.Data.Country}}{{else}}Unknown{{end}}',
|
||||
ipAddress: '{{.Data.IPAddress}}',
|
||||
device: '{{.Data.Device}}',
|
||||
method: '{{.Data.Method}}',
|
||||
dateTime: '{{.Data.DateTime.Format "January 2, 2006 at 3:04 PM MST"}}'
|
||||
}
|
||||
};
|
||||
@@ -58,6 +61,7 @@ NewSignInEmail.PreviewProps = {
|
||||
location: 'San Francisco, USA',
|
||||
ipAddress: '203.0.113.42',
|
||||
device: 'Chrome on macOS',
|
||||
method: 'Passkey',
|
||||
dateTime: 'January 2, 2026 at 3:04 PM UTC'
|
||||
}
|
||||
};
|
||||
|
||||
@@ -202,6 +202,14 @@
|
||||
"this_can_be_useful_for_selfsigned_certificates": "This can be useful for self-signed certificates.",
|
||||
"enabled_emails": "Enabled Emails",
|
||||
"email_login_notification": "Email Login Notification",
|
||||
"email_login_notification_description": "Choose when users receive an email after signing in.",
|
||||
"login_notification_disabled_description": "Do not send sign-in notification emails or use browser recognition cookies.",
|
||||
"login_notification_always": "Every sign-in",
|
||||
"login_notification_always_description": "Send an email after every successful sign-in. No browser recognition cookie is used.",
|
||||
"login_notification_ip_and_user_agent": "New IP address or browser",
|
||||
"login_notification_ip_and_user_agent_description": "Send an email when the exact combination of IP address and browser User-Agent is not in the user's sign-in history. No browser recognition cookie is used.",
|
||||
"login_notification_browser_recognition": "Unrecognized browser",
|
||||
"login_notification_browser_recognition_description": "Send an email unless the browser is recognized by its cookie or a matching IP address and User-Agent in the user's sign-in history.",
|
||||
"send_an_email_to_the_user_when_they_log_in_from_a_new_device": "Send an email to the user when they log in from a new device.",
|
||||
"emai_login_code_requested_by_user": "Email Login Code Requested by User",
|
||||
"allow_users_to_sign_in_with_a_login_code_sent_to_their_email": "Allows users to bypass passkeys by requesting a login code sent to their email. This significantly reduces security as anyone with access to the user's email can gain entry.",
|
||||
|
||||
@@ -31,7 +31,7 @@ export type AllAppConfig = AppConfig & {
|
||||
smtpPassword: string;
|
||||
smtpTls: 'none' | 'starttls' | 'tls';
|
||||
smtpSkipCertVerify: boolean;
|
||||
emailLoginNotificationEnabled: boolean;
|
||||
emailLoginNotificationMode: 'disabled' | 'always' | 'ipAndUserAgent' | 'browserRecognition';
|
||||
emailApiKeyExpirationEnabled: boolean;
|
||||
// LDAP
|
||||
ldapUrl: string;
|
||||
|
||||
+55
-9
@@ -30,6 +30,22 @@
|
||||
tls: 'TLS'
|
||||
};
|
||||
|
||||
const notificationOptions = $derived({
|
||||
disabled: { label: m.never(), description: m.login_notification_disabled_description() },
|
||||
always: {
|
||||
label: m.login_notification_always(),
|
||||
description: m.login_notification_always_description()
|
||||
},
|
||||
ipAndUserAgent: {
|
||||
label: m.login_notification_ip_and_user_agent(),
|
||||
description: m.login_notification_ip_and_user_agent_description()
|
||||
},
|
||||
browserRecognition: {
|
||||
label: m.login_notification_browser_recognition(),
|
||||
description: m.login_notification_browser_recognition_description()
|
||||
}
|
||||
});
|
||||
|
||||
let isSendingTestEmail = $state(false);
|
||||
|
||||
const formSchema = z
|
||||
@@ -49,7 +65,12 @@
|
||||
emailOneTimeAccessAsUnauthenticatedEnabled: z.boolean(),
|
||||
emailVerificationEnabled: z.boolean(),
|
||||
emailOneTimeAccessAsAdminEnabled: z.boolean(),
|
||||
emailLoginNotificationEnabled: z.boolean(),
|
||||
emailLoginNotificationMode: z.enum([
|
||||
'disabled',
|
||||
'always',
|
||||
'ipAndUserAgent',
|
||||
'browserRecognition'
|
||||
]),
|
||||
emailApiKeyExpirationEnabled: z.boolean()
|
||||
})
|
||||
.superRefine((data, ctx) => {
|
||||
@@ -63,7 +84,6 @@
|
||||
'emailOneTimeAccessAsUnauthenticatedEnabled',
|
||||
'emailVerificationEnabled',
|
||||
'emailOneTimeAccessAsAdminEnabled',
|
||||
'emailLoginNotificationEnabled',
|
||||
'emailApiKeyExpirationEnabled'
|
||||
];
|
||||
|
||||
@@ -84,7 +104,8 @@
|
||||
const anyProvided = requiredSmtpFields.some((f) => !!data[f]);
|
||||
requireFieldsWhen(anyProvided, m.smtp_field_required_when_other_provided());
|
||||
|
||||
const emailEnabled = emailFields.some((f) => data[f]);
|
||||
const emailEnabled =
|
||||
data.emailLoginNotificationMode !== 'disabled' || emailFields.some((f) => data[f]);
|
||||
requireFieldsWhen(emailEnabled, m.smtp_field_required_when_email_enabled());
|
||||
});
|
||||
|
||||
@@ -194,12 +215,37 @@
|
||||
</div>
|
||||
<h4 class="mt-10 text-lg font-semibold">{m.enabled_emails()}</h4>
|
||||
<div class="mt-4 flex flex-col gap-5">
|
||||
<SwitchWithLabel
|
||||
id="email-login-notification"
|
||||
label={m.email_login_notification()}
|
||||
description={m.send_an_email_to_the_user_when_they_log_in_from_a_new_device()}
|
||||
bind:checked={$inputs.emailLoginNotificationEnabled.value}
|
||||
/>
|
||||
<Field.Field>
|
||||
<div>
|
||||
<Field.Label for="email-login-notification">{m.email_login_notification()}</Field.Label>
|
||||
<Field.Description>{m.email_login_notification_description()}</Field.Description>
|
||||
</div>
|
||||
<Select.Root
|
||||
type="single"
|
||||
disabled={$appConfigStore.uiConfigDisabled}
|
||||
value={$inputs.emailLoginNotificationMode.value}
|
||||
allowDeselect={false}
|
||||
onValueChange={(value) =>
|
||||
($inputs.emailLoginNotificationMode.value =
|
||||
value as AllAppConfig['emailLoginNotificationMode'])}
|
||||
>
|
||||
<Select.Trigger id="email-login-notification" class="w-full">
|
||||
{notificationOptions[$inputs.emailLoginNotificationMode.value].label}
|
||||
</Select.Trigger>
|
||||
<Select.Content class="w-[calc(var(--bits-select-anchor-width)+--spacing(3))]">
|
||||
<Select.Group>
|
||||
{#each Object.entries(notificationOptions) as [value, option] (value)}
|
||||
<Select.Item {value} label={option.label}>
|
||||
<div class="flex flex-col items-start gap-1 whitespace-normal">
|
||||
<span class="font-medium">{option.label}</span>
|
||||
<span class="text-muted-foreground text-xs">{option.description}</span>
|
||||
</div>
|
||||
</Select.Item>
|
||||
{/each}
|
||||
</Select.Group>
|
||||
</Select.Content>
|
||||
</Select.Root>
|
||||
</Field.Field>
|
||||
<SwitchWithLabel
|
||||
id="email-verification"
|
||||
label={m.email_verification()}
|
||||
|
||||
@@ -246,6 +246,7 @@ test('Update email configuration', async ({ page }) => {
|
||||
await page.getByLabel('SMTP Password').fill('password');
|
||||
await page.getByLabel('SMTP From').fill('test@gmail.com');
|
||||
await page.getByLabel('Email Login Notification').click();
|
||||
await page.getByRole('option', { name: /^Unrecognized browser/ }).click();
|
||||
await page.getByLabel('Email Login Code Requested by User').click();
|
||||
await page.getByLabel('Email Login Code from Admin').click();
|
||||
await page.getByLabel('API Key Expiration').click();
|
||||
@@ -259,7 +260,7 @@ test('Update email configuration', async ({ page }) => {
|
||||
await expect(page.getByLabel('SMTP User')).toHaveValue('test@gmail.com');
|
||||
await expect(page.getByLabel('SMTP Password')).toHaveValue('password');
|
||||
await expect(page.getByLabel('SMTP From')).toHaveValue('test@gmail.com');
|
||||
await expect(page.getByLabel('Email Login Notification')).toBeChecked();
|
||||
await expect(page.getByLabel('Email Login Notification')).toContainText('Unrecognized browser');
|
||||
await expect(page.getByLabel('Email Login Code Requested by User')).toBeChecked();
|
||||
await expect(page.getByLabel('Email Login Code from Admin')).toBeChecked();
|
||||
await expect(page.getByLabel('API Key Expiration')).toBeChecked();
|
||||
|
||||
Reference in New Issue
Block a user