Files
probo/pkg/iam/auth_service.go
Bryan Frimin 11770b4058 Add OAuth2/OpenID Connect authorization server
Implement a full OAuth2 2.0 and OpenID Connect 1.0 authorization
server with support for authorization code flow (with PKCE),
refresh token rotation, device authorization grant, dynamic
client registration, token introspection, and token revocation.

Includes database schema, coredata layer, service logic, HTTP
handlers, OIDC discovery endpoint, and JWKS publishing.

Signed-off-by: Bryan Frimin <bryan@getprobo.com>
2026-04-19 12:00:53 +02:00

720 lines
19 KiB
Go

// Copyright (c) 2025-2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package iam
import (
"context"
"errors"
"fmt"
"time"
"go.gearno.de/kit/pg"
"go.probo.inc/probo/packages/emails"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/crypto/hash"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/mail"
"go.probo.inc/probo/pkg/statelesstoken"
"go.probo.inc/probo/pkg/validator"
)
type (
AuthService struct {
*Service
}
ResetPasswordRequest struct {
Token string
Password string
}
ChangePasswordRequest struct {
CurrentPassword string
NewPassword string
}
ActivateAccountRequest struct {
InvitationToken string
}
CreateIdentityWithPasswordRequest struct {
Email mail.Addr
Password string
FullName string
}
SendMagicLinkRequest struct {
Email mail.Addr
URLPath string
OrganizationID gid.GID
Continue *string
// If users tries to connect to compliance page, we must brand the emails accordingly
CompliancePageID *gid.GID
}
PasswordResetData struct {
Email mail.Addr `json:"email"`
}
MagicLinkData struct {
Email mail.Addr `json:"email"`
Continue *string `json:"continue"`
}
)
const (
TokenTypeOrganizationInvitation = "organization_invitation"
TokenTypePasswordReset = "password_reset"
TokenTypeMagicLink = "magic_link"
)
func NewAuthService(svc *Service) *AuthService {
return &AuthService{Service: svc}
}
func (req ActivateAccountRequest) Validate() error {
v := validator.New()
v.Check(req.InvitationToken, "invitationToken", validator.NotEmpty())
return v.Error()
}
func (req ResetPasswordRequest) Validate() error {
v := validator.New()
v.Check(req.Token, "token", validator.NotEmpty())
v.Check(req.Password, "password", PasswordValidator())
return v.Error()
}
func (req ChangePasswordRequest) Validate() error {
v := validator.New()
// We cannot use PasswordValidator here because legacy password may not be aligned with the current password
// policy, therefore we at least enforce a maximum length to mitigate DDoS attacks.
v.Check(req.CurrentPassword, "currentPassword", validator.NotEmpty(), validator.MaxLen(255))
v.Check(req.NewPassword, "newPassword", PasswordValidator())
return v.Error()
}
func (req CreateIdentityWithPasswordRequest) Validate() error {
v := validator.New()
v.Check(req.FullName, "fullName", validator.NotEmpty(), validator.MinLen(1), validator.MaxLen(255))
v.Check(req.Password, "password", PasswordValidator())
return v.Error()
}
func (s *AuthService) ActivateAccount(
ctx context.Context,
req *ActivateAccountRequest,
) (*coredata.Identity, *coredata.MembershipProfile, error) {
if err := req.Validate(); err != nil {
return nil, nil, fmt.Errorf("invalid request: %w", err)
}
payload, err := statelesstoken.ValidateToken[InvitationTokenData](s.tokenSecret, TokenTypeOrganizationInvitation, req.InvitationToken)
if err != nil {
return nil, nil, NewInvalidTokenError()
}
var (
scope = coredata.NewScopeFromObjectID(payload.Data.InvitationID)
invitation = &coredata.Invitation{}
profile *coredata.MembershipProfile
identity *coredata.Identity
now = time.Now()
)
if err = s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
err := invitation.LoadByID(ctx, tx, scope, payload.Data.InvitationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
return NewInvitationNotFoundError(payload.Data.InvitationID)
}
return fmt.Errorf("cannot load invitation: %w", err)
}
if invitation.AcceptedAt != nil {
return NewInvitationAlreadyAcceptedError(payload.Data.InvitationID)
}
if invitation.ExpiresAt.Before(now) {
return NewInvitationExpiredError(payload.Data.InvitationID)
}
profile = &coredata.MembershipProfile{}
if err := profile.LoadByID(ctx, tx, scope, invitation.UserID); err != nil {
return fmt.Errorf("cannot load user: %w", err)
}
if profile.Source == coredata.ProfileSourceSCIM {
return NewUserManagedBySCIMError(profile.ID)
}
if profile.State == coredata.ProfileStateInactive {
profile.State = coredata.ProfileStateActive
profile.UpdatedAt = now
if err := profile.Update(ctx, tx, scope); err != nil {
return fmt.Errorf("cannot update user: %w", err)
}
}
identity = &coredata.Identity{}
if err := identity.LoadByID(ctx, tx, profile.IdentityID); err != nil {
return fmt.Errorf("cannot load identity: %w", err)
}
identity.EmailAddressVerified = true
identity.UpdatedAt = now
err = identity.Update(ctx, tx)
if err != nil {
return fmt.Errorf("cannot update identity: %w", err)
}
invitation.AcceptedAt = &now
if err := invitation.Update(ctx, tx, scope); err != nil {
if err == coredata.ErrResourceNotFound {
return NewInvitationNotFoundError(payload.Data.InvitationID)
}
return fmt.Errorf("cannot update invitation: %w", err)
}
// Expire other pending invitations for user
invitations := &coredata.Invitations{}
onlyPending := coredata.NewInvitationFilter([]coredata.InvitationStatus{coredata.InvitationStatusPending})
if err := invitations.ExpireByUserID(
ctx,
tx,
coredata.NewScopeFromObjectID(invitation.OrganizationID),
invitation.UserID,
onlyPending,
); err != nil {
return fmt.Errorf("cannot expire pending invitations: %w", err)
}
return nil
},
); err != nil {
return nil, nil, err
}
return identity, profile, nil
}
func (s AuthService) GetResetPasswordToken(ctx context.Context, email mail.Addr) (string, error) {
token, err := statelesstoken.NewToken(
s.tokenSecret,
TokenTypePasswordReset,
s.passwordResetTokenValidity,
PasswordResetData{Email: email},
)
if err != nil {
return "", fmt.Errorf("cannot generate password create token: %w", err)
}
return token, nil
}
func (s AuthService) ResetPassword(
ctx context.Context,
req *ResetPasswordRequest,
) error {
if err := req.Validate(); err != nil {
return fmt.Errorf("invalid request: %w", err)
}
payload, err := statelesstoken.ValidateToken[PasswordResetData](s.tokenSecret, TokenTypePasswordReset, req.Token)
if err != nil {
return NewInvalidTokenError()
}
hashedPassword, err := s.hp.HashPassword([]byte(req.Password))
if err != nil {
return fmt.Errorf("cannot hash password: %w", err)
}
return s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
err := identity.LoadByEmail(ctx, tx, payload.Data.Email)
if err != nil {
if err == coredata.ErrResourceNotFound {
return nil // Don't leak information about non-existent identities
}
return fmt.Errorf("cannot load identity: %w", err)
}
identity.HashedPassword = hashedPassword
identity.UpdatedAt = time.Now()
err = identity.Update(ctx, tx)
if err != nil {
if err == coredata.ErrResourceNotFound {
return nil // Don't leak information about non-existent identities
}
return fmt.Errorf("cannot update identity: %w", err)
}
return nil
},
)
}
func (s AuthService) SendPasswordResetInstructionByEmail(
ctx context.Context,
email mail.Addr,
) error {
token, err := s.GetResetPasswordToken(ctx, email)
if err != nil {
return fmt.Errorf("cannot generate password reset token: %w", err)
}
return s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
if err := identity.LoadByEmail(ctx, tx, email); err != nil {
if err == coredata.ErrResourceNotFound {
return nil // Don't leak information about non-existent identities
}
return fmt.Errorf("cannot load identity: %w", err)
}
emailPresenter := emails.NewPresenter(s.fm, s.bucket, s.baseURL, identity.FullName)
subject, textBody, htmlBody, err := emailPresenter.RenderPasswordReset(
ctx,
"/auth/reset-password",
token,
)
if err != nil {
return fmt.Errorf("cannot render password reset email: %w", err)
}
passwordResetEmail := coredata.NewEmail(
identity.FullName,
identity.EmailAddress,
subject,
textBody,
htmlBody,
nil,
)
err = passwordResetEmail.Insert(ctx, tx)
if err != nil {
return fmt.Errorf("cannot insert email: %w", err)
}
return nil
},
)
}
func (s AuthService) CreateIdentityWithPassword(
ctx context.Context,
req *CreateIdentityWithPasswordRequest,
) (*coredata.Identity, *coredata.Session, error) {
if s.disableSignup { // TODO Rename this one to disableSignup
return nil, nil, NewErrSignupDisabled()
}
if err := req.Validate(); err != nil {
return nil, nil, fmt.Errorf("invalid request: %w", err)
}
hashedPassword, err := s.hp.HashPassword([]byte(req.Password))
if err != nil {
return nil, nil, fmt.Errorf("cannot hash password: %w", err)
}
var (
now = time.Now()
identity = &coredata.Identity{
ID: gid.New(gid.NilTenant, coredata.IdentityEntityType),
EmailAddress: req.Email,
FullName: req.FullName,
HashedPassword: hashedPassword,
EmailAddressVerified: false,
CreatedAt: now,
UpdatedAt: now,
}
session = coredata.NewRootSession(identity.ID, coredata.AuthMethodPassword, 24*time.Hour*7)
)
confirmationToken, err := statelesstoken.NewToken(
s.tokenSecret,
TokenTypeEmailConfirmation,
24*time.Hour,
EmailConfirmationData{IdentityID: identity.ID, Email: identity.EmailAddress},
)
if err != nil {
return nil, nil, fmt.Errorf("cannot generate confirmation token: %w", err)
}
emailPresenter := emails.NewPresenter(s.fm, s.bucket, s.baseURL, req.FullName)
subject, textBody, htmlBody, err := emailPresenter.RenderConfirmEmail(ctx, "/auth/verify-email", confirmationToken)
if err != nil {
return nil, nil, fmt.Errorf("cannot render confirmation email: %w", err)
}
confirmationEmail := coredata.NewEmail(
req.FullName,
identity.EmailAddress,
subject,
textBody,
htmlBody,
nil,
)
err = s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
err := identity.Insert(ctx, tx)
if err != nil {
if err == coredata.ErrResourceAlreadyExists {
return NewIdentityAlreadyExistsError(identity.EmailAddress)
}
return fmt.Errorf("cannot insert identity: %w", err)
}
if err := confirmationEmail.Insert(ctx, tx); err != nil {
return fmt.Errorf("cannot insert email: %w", err)
}
if err := session.Insert(ctx, tx); err != nil {
return fmt.Errorf("cannot insert session: %w", err)
}
return nil
},
)
return identity, session, err
}
func (s AuthService) OpenSessionWithSAML(ctx context.Context, identityID gid.GID) (*coredata.Session, error) {
session := &coredata.Session{}
err := s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) (err error) {
session = coredata.NewRootSession(identityID, coredata.AuthMethodSAML, s.sessionDuration)
err = session.Insert(ctx, conn)
if err != nil {
return fmt.Errorf("cannot insert session: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return session, nil
}
func (s AuthService) OpenSessionWithOIDC(ctx context.Context, identityID gid.GID, authMethod coredata.AuthMethod) (*coredata.Session, error) {
session := &coredata.Session{}
err := s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) (err error) {
session = coredata.NewRootSession(identityID, authMethod, s.sessionDuration)
err = session.Insert(ctx, conn)
if err != nil {
return fmt.Errorf("cannot insert session: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return session, nil
}
func (s AuthService) CheckCredentials(
ctx context.Context,
email mail.Addr,
password string,
) (*coredata.Identity, error) {
v := validator.New()
v.Check(password, "password", PasswordValidator())
err := v.Error()
if err != nil {
return nil, NewInvalidPasswordError("invalid password")
}
identity := &coredata.Identity{}
err = s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) error {
err := identity.LoadByEmail(ctx, conn, email)
if err != nil {
// Do not leak information about non-existent identities
if err != coredata.ErrResourceNotFound {
return fmt.Errorf("cannot load identity by email: %w", err)
}
}
// Perform a password comparison even when the identity does not exist to mitigate timing attacks
// and prevent revealing account existence.
if identity.ID == gid.Nil {
_, _ = s.hp.ComparePasswordAndHash([]byte(password+"qwertyuiop1234567890"), []byte("qwertyuiop1234567890"))
return NewInvalidCredentialsError("invalid email or password")
}
isPasswordMatch, err := s.hp.ComparePasswordAndHash([]byte(password), identity.HashedPassword)
if err != nil {
return fmt.Errorf("cannot verify password: %w", err)
}
if !isPasswordMatch {
return NewInvalidCredentialsError("invalid email or password")
}
return nil
},
)
return identity, err
}
func (s AuthService) OpenSessionWithPassword(ctx context.Context, identityID gid.GID) (*coredata.Session, error) {
session := &coredata.Session{}
err := s.pg.WithTx(
ctx,
func(ctx context.Context, conn pg.Tx) (err error) {
session = coredata.NewRootSession(identityID, coredata.AuthMethodPassword, s.sessionDuration)
err = session.Insert(ctx, conn)
if err != nil {
return fmt.Errorf("cannot insert session: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return session, nil
}
func (s AuthService) SendMagicLink(ctx context.Context, req *SendMagicLinkRequest) error {
tokenString, err := statelesstoken.NewToken(
s.tokenSecret,
TokenTypeMagicLink,
s.magicLinkTokenValidity,
MagicLinkData{
Email: req.Email,
Continue: req.Continue,
},
)
if err != nil {
return fmt.Errorf("cannot generate magic link token: %w", err)
}
return s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
hashedToken := HashToken(tokenString)
token := &coredata.Token{
ID: gid.New(gid.NilTenant, coredata.TokenEntityType),
HashedValue: hashedToken,
CreatedAt: time.Now(),
}
if err := token.Insert(ctx, tx); err != nil {
return fmt.Errorf("cannot insert token: %w", err)
}
fullName := req.Email.Username()
identity := &coredata.Identity{}
organization := &coredata.Organization{}
if err := identity.LoadByEmail(ctx, tx, req.Email); err == nil {
if identity.FullName != "" {
fullName = identity.FullName
}
} else {
if !errors.Is(err, coredata.ErrResourceNotFound) {
return fmt.Errorf("cannot load identity: %w", err)
}
}
if err := organization.LoadByID(ctx, tx, coredata.NewNoScope(), req.OrganizationID); err != nil {
return fmt.Errorf("cannot load organization: %w", err)
}
emailPresenterCfg := emails.DefaultPresenterConfig(s.bucket, s.baseURL)
if req.CompliancePageID != nil {
var err error
emailPresenterCfg, err = s.CompliancePageService.EmailPresenterConfig(ctx, *req.CompliancePageID)
if err != nil {
return fmt.Errorf("cannot get compliance page email presenter config: %w", err)
}
}
emailPresenter := emails.NewPresenterFromConfig(s.fm, emailPresenterCfg, fullName)
subject, textBody, htmlBody, err := emailPresenter.RenderMagicLink(
ctx,
req.URLPath,
tokenString,
s.magicLinkTokenValidity,
organization.Name,
)
if err != nil {
return fmt.Errorf("cannot render magic link email: %w", err)
}
var emailOpts *coredata.EmailOptions
if req.CompliancePageID != nil {
emailOpts = &coredata.EmailOptions{
SenderName: new(organization.Name),
}
}
magicLinkEmail := coredata.NewEmail(
fullName,
req.Email,
subject,
textBody,
htmlBody,
emailOpts,
)
if err := magicLinkEmail.Insert(ctx, tx); err != nil {
return fmt.Errorf("cannot insert email: %w", err)
}
return nil
},
)
}
func (s AuthService) GetMagicLinkEmail(ctx context.Context, tokenString string) (mail.Addr, error) {
payload, err := statelesstoken.ValidateToken[MagicLinkData](s.tokenSecret, TokenTypeMagicLink, tokenString)
if err != nil {
var errExpired *statelesstoken.ErrExpiredToken
if errors.As(err, &errExpired) {
return mail.Nil, NewExpiredTokenError()
}
return mail.Nil, NewInvalidTokenError()
}
return payload.Data.Email, nil
}
func (s AuthService) OpenSessionWithMagicLink(ctx context.Context, tokenString string) (*coredata.Identity, *coredata.Session, *string, error) {
var (
now = time.Now()
session = &coredata.Session{}
identity = &coredata.Identity{}
)
payload, err := statelesstoken.ValidateToken[MagicLinkData](s.tokenSecret, TokenTypeMagicLink, tokenString)
if err != nil {
var errExpired *statelesstoken.ErrExpiredToken
if errors.As(err, &errExpired) {
return nil, nil, nil, NewExpiredTokenError()
}
return nil, nil, nil, NewInvalidTokenError()
}
if err := s.pg.WithTx(
ctx,
func(ctx context.Context, tx pg.Tx) error {
hashedValue := HashToken(tokenString)
token := &coredata.Token{}
if err := token.LoadByHashedValueForUpdate(ctx, tx, hashedValue); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return NewInvalidTokenError()
}
return fmt.Errorf("cannot load token by hashed value: %w", err)
}
err := identity.LoadByEmail(ctx, tx, payload.Data.Email)
if err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
identity = &coredata.Identity{
ID: gid.New(gid.NilTenant, coredata.IdentityEntityType),
EmailAddress: payload.Data.Email,
EmailAddressVerified: true,
CreatedAt: now,
UpdatedAt: now,
}
if err := identity.Insert(ctx, tx); err != nil {
return fmt.Errorf("cannot create identity: %w", err)
}
} else {
return fmt.Errorf("cannot load identity by email: %w", err)
}
}
session = coredata.NewRootSession(identity.ID, coredata.AuthMethodMagicLink, s.sessionDuration)
err = session.Insert(ctx, tx)
if err != nil {
return fmt.Errorf("cannot insert session: %w", err)
}
if err := token.Delete(ctx, tx); err != nil {
return fmt.Errorf("cannot delete token: %w", err)
}
return nil
},
); err != nil {
return nil, nil, nil, err
}
return identity, session, payload.Data.Continue, nil
}
func HashToken(token string) []byte {
return hash.SHA256String(token)
}