// Copyright (c) 2025-2026 Probo Inc . // // Permission is hereby granted, free of charge, to any person obtaining a copy // of this software and associated documentation files (the "Software"), to deal // in the Software without restriction, including without limitation the rights // to use, copy, modify, merge, publish, distribute, sublicense, and/or sell // copies of the Software, and to permit persons to whom the Software is // furnished to do so, subject to the following conditions: // // The above copyright notice and this permission notice shall be included in // all copies or substantial portions of the Software. // // THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR // IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, // FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE // AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER // LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, // OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE // 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/gid" "go.probo.inc/probo/pkg/mail" "go.probo.inc/probo/pkg/page" "go.probo.inc/probo/pkg/securetoken" "go.probo.inc/probo/pkg/statelesstoken" "go.probo.inc/probo/pkg/validator" ) type ( AccountService struct { *Service } PersonalAPIKeyTokenData struct { Version int `json:"v"` KeyID gid.GID `json:"kid"` PrincipalID gid.GID `json:"pid"` IssuedAt time.Time `json:"iat"` } EmailConfirmationData struct { IdentityID gid.GID `json:"uid"` Email mail.Addr `json:"email"` } ChangeEmailRequest struct { NewEmail mail.Addr Password string } UpdateIdentityRequest struct { FullName string `json:"fullName"` } UpdateLocaleRequest struct { Locale string `json:"locale"` } ) // SupportedIdentityLocales are short URL locale tags accepted for // Identity.locale. Keep in sync with the compliance-portal URL_LOCALES list. var SupportedIdentityLocales = []string{ "en", "fr", "de", "es", "id", "it", "ja", "ko", "pl", "pt", "tr", "uk", "zh", } const ( TokenTypeEmailConfirmation = "email_confirmation" ) func NewAccountService(svc *Service) *AccountService { return &AccountService{Service: svc} } func (req ChangeEmailRequest) Validate() error { v := validator.New() v.Check(req.Password, "password", validator.NotEmpty(), validator.MaxLen(255)) // 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. return v.Error() } func (req UpdateIdentityRequest) Validate() error { v := validator.New() v.Check(req.FullName, "full_name", validator.NotEmpty(), validator.MinLen(2), validator.MaxLen(255)) return v.Error() } func (req UpdateLocaleRequest) Validate() error { v := validator.New() v.Check( req.Locale, "locale", validator.NotEmpty(), validator.MaxLen(8), validator.OneOfSlice(SupportedIdentityLocales), ) return v.Error() } func (s AccountService) ChangeEmail(ctx context.Context, identityID gid.GID, req *ChangeEmailRequest) error { if err := req.Validate(); err != nil { return fmt.Errorf("invalid request: %w", err) } confirmationToken, err := statelesstoken.NewToken( s.tokenSecret, TokenTypeEmailConfirmation, s.emailConfirmationTokenValidity, EmailConfirmationData{IdentityID: identityID, Email: req.NewEmail}, ) if err != nil { return fmt.Errorf("cannot generate confirmation token: %w", err) } return s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { identity := &coredata.Identity{} err := identity.LoadByID(ctx, tx, identityID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(identityID) } return fmt.Errorf("cannot load identity: %w", err) } isPasswordMatch, err := s.hp.ComparePasswordAndHash([]byte(req.Password), identity.HashedPassword) if err != nil { return fmt.Errorf("cannot compare password: %w", err) } if !isPasswordMatch { return NewInvalidPasswordError("invalid password") } identity.EmailAddress = req.NewEmail identity.EmailAddressVerified = false identity.UpdatedAt = time.Now() err = identity.Update(ctx, tx) if err != nil { return fmt.Errorf("cannot update identity: %w", err) } emailPresenter := emails.NewPresenter(s.baseURL, identity.FullName) subject, textBody, htmlBody, err := emailPresenter.RenderConfirmEmail(ctx, "/auth/verify-email", confirmationToken) if err != nil { return fmt.Errorf("cannot render confirmation email: %w", err) } confirmationEmail := coredata.NewEmail( identity.FullName, identity.EmailAddress, subject, textBody, htmlBody, nil, ) err = confirmationEmail.Insert(ctx, tx) if err != nil { return fmt.Errorf("cannot insert confirmation email: %w", err) } return nil }, ) } func (s AccountService) VerifyEmail(ctx context.Context, token string) error { payload, err := statelesstoken.ValidateToken[EmailConfirmationData](s.tokenSecret, TokenTypeEmailConfirmation, token) if err != nil { return NewInvalidTokenError() } return s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { identity := &coredata.Identity{} err := identity.LoadByID(ctx, tx, payload.Data.IdentityID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(payload.Data.IdentityID) } return fmt.Errorf("cannot load identity: %w", err) } if identity.EmailAddress != payload.Data.Email { return NewEmailVerificationMismatchError() } if identity.EmailAddressVerified { return NewEmailAlreadyVerifiedError() } identity.EmailAddressVerified = true identity.UpdatedAt = time.Now() err = identity.Update(ctx, tx) if err != nil { return fmt.Errorf("cannot update identity: %w", err) } return nil }, ) } func (s AccountService) ResendVerificationEmail(ctx context.Context, email mail.Addr) error { 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) } if identity.EmailAddressVerified { return nil // Don't leak information about already-verified identities } confirmationToken, err := statelesstoken.NewToken( s.tokenSecret, TokenTypeEmailConfirmation, s.emailConfirmationTokenValidity, EmailConfirmationData{IdentityID: identity.ID, Email: identity.EmailAddress}, ) if err != nil { return fmt.Errorf("cannot generate confirmation token: %w", err) } emailPresenter := emails.NewPresenter(s.baseURL, identity.FullName) subject, textBody, htmlBody, err := emailPresenter.RenderConfirmEmail(ctx, "/auth/verify-email", confirmationToken) if err != nil { return fmt.Errorf("cannot render confirmation email: %w", err) } confirmationEmail := coredata.NewEmail( identity.FullName, identity.EmailAddress, subject, textBody, htmlBody, nil, ) if err := confirmationEmail.Insert(ctx, tx); err != nil { return fmt.Errorf("cannot insert confirmation email: %w", err) } return nil }, ) } func (s *AccountService) ListPendingInvitations( ctx context.Context, userID gid.GID, cursor *page.Cursor[coredata.InvitationOrderField], ) (*page.Page[*coredata.Invitation, coredata.InvitationOrderField], error) { var ( scope = coredata.NewScopeFromObjectID(userID) invitations coredata.Invitations ) err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { profile := coredata.MembershipProfile{} err := profile.LoadByID(ctx, conn, scope, userID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(userID) } return fmt.Errorf("cannot load identity: %w", err) } onlyPending := coredata.NewInvitationFilter([]coredata.InvitationStatus{coredata.InvitationStatusPending}) err = invitations.LoadByUserID(ctx, conn, scope, userID, cursor, onlyPending) if err != nil { return fmt.Errorf("cannot load invitations: %w", err) } return nil }, ) if err != nil { return nil, err } return page.NewPage(invitations, cursor), nil } func (s AccountService) ChangePassword(ctx context.Context, identityID gid.GID, currentSessionID gid.GID, req *ChangePasswordRequest) error { if err := req.Validate(); err != nil { return fmt.Errorf("invalid request: %w", err) } return s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { identity := &coredata.Identity{} err := identity.LoadByID(ctx, tx, identityID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(identityID) } return fmt.Errorf("cannot load identity: %w", err) } isLegacyPasswordMatch, err := s.hp.ComparePasswordAndHash([]byte(req.CurrentPassword), identity.HashedPassword) if err != nil { return fmt.Errorf("cannot compare legacy password: %w", err) } if !isLegacyPasswordMatch { return NewInvalidPasswordError("invalid current password") } newPasswordHash, err := s.hp.HashPassword([]byte(req.NewPassword)) if err != nil { return fmt.Errorf("cannot hash new password: %w", err) } identity.HashedPassword = newPasswordHash identity.UpdatedAt = time.Now() err = identity.Update(ctx, tx) if err != nil { return fmt.Errorf("cannot update identity: %w", err) } sessions := coredata.Sessions{} if _, err := sessions.ExpireAllForIdentityExceptOneSession(ctx, tx, identity.ID, currentSessionID); err != nil { return fmt.Errorf("cannot expire other sessions: %w", err) } // TODO: email to notify identity that their password has been changed return nil }, ) } func (s AccountService) CountSessions(ctx context.Context, identityID gid.GID) (int, error) { var count int err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) (err error) { sessions := coredata.Sessions{} count, err = sessions.CountByIdentityID(ctx, conn, identityID) if err != nil { return fmt.Errorf("cannot count sessions: %w", err) } return nil }, ) return count, err } func (s AccountService) ListSessions( ctx context.Context, identityID gid.GID, cursor *page.Cursor[coredata.SessionOrderField], ) (*page.Page[*coredata.Session, coredata.SessionOrderField], error) { var sessions coredata.Sessions err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { err := sessions.LoadByIdentityID(ctx, conn, identityID, cursor) if err != nil { return fmt.Errorf("cannot load sessions: %w", err) } return nil }, ) if err != nil { return nil, err } return page.NewPage(sessions, cursor), nil } func (s AccountService) GetIdentity(ctx context.Context, identityID gid.GID) (*coredata.Identity, error) { identity := &coredata.Identity{} err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { err := identity.LoadByID(ctx, conn, identityID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(identityID) } return fmt.Errorf("cannot load identity: %w", err) } return nil }, ) if err != nil { return nil, err } return identity, nil } func (s AccountService) UpdateIdentity(ctx context.Context, identityID gid.GID, req *UpdateIdentityRequest) (*coredata.Identity, error) { if err := req.Validate(); err != nil { return nil, fmt.Errorf("invalid request: %w", err) } identity := &coredata.Identity{} err := s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { err := identity.LoadByID(ctx, tx, identityID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(identityID) } return fmt.Errorf("cannot load identity: %w", err) } identity.FullName = req.FullName identity.UpdatedAt = time.Now() if err := identity.Update(ctx, tx); err != nil { return fmt.Errorf("cannot update identity: %w", err) } return nil }, ) if err != nil { return nil, err } return identity, nil } func (s AccountService) UpdateLocale(ctx context.Context, identityID gid.GID, req *UpdateLocaleRequest) (*coredata.Identity, error) { if err := req.Validate(); err != nil { return nil, fmt.Errorf("invalid request: %w", err) } identity := &coredata.Identity{} err := s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { err := identity.LoadByID(ctx, tx, identityID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(identityID) } return fmt.Errorf("cannot load identity: %w", err) } identity.Locale = &req.Locale identity.UpdatedAt = time.Now() if err := identity.Update(ctx, tx); err != nil { return fmt.Errorf("cannot update identity locale: %w", err) } return nil }, ) if err != nil { return nil, err } return identity, nil } func (s AccountService) ListPersonalAPIKeys( ctx context.Context, identityID gid.GID, cursor *page.Cursor[coredata.PersonalAPIKeyOrderField], ) (*page.Page[*coredata.PersonalAPIKey, coredata.PersonalAPIKeyOrderField], error) { var personalAccessTokens coredata.PersonalAPIKeys err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { err := personalAccessTokens.LoadByIdentityID(ctx, conn, identityID) if err != nil { return fmt.Errorf("cannot load personal access tokens: %w", err) } return nil }, ) if err != nil { return nil, err } return page.NewPage(personalAccessTokens, cursor), nil } func (s AccountService) CountPersonalAPIKeys(ctx context.Context, identityID gid.GID) (int, error) { var count int err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) (err error) { personalAccessTokens := coredata.PersonalAPIKeys{} count, err = personalAccessTokens.CountByIdentityID(ctx, conn, identityID) if err != nil { return fmt.Errorf("cannot count personal access tokens: %w", err) } return nil }, ) return count, err } func (s *AccountService) RevealPersonalAPIKeyToken( ctx context.Context, identityID gid.GID, personalAPIKeyID gid.GID, ) (string, error) { var token string err := s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) (err error) { personalAPIKey := &coredata.PersonalAPIKey{} if err := personalAPIKey.LoadByID(ctx, tx, personalAPIKeyID); err != nil { if err == coredata.ErrResourceNotFound { return NewPersonalAPIKeyNotFoundError(personalAPIKeyID) } return fmt.Errorf("cannot load personal api key: %w", err) } if personalAPIKey.IdentityID != identityID { return NewPersonalAPIKeyNotFoundError(personalAPIKeyID) } token, err = securetoken.Sign( personalAPIKey.ID.String(), s.tokenSecret, ) if err != nil { return fmt.Errorf("cannot generate personal api key token: %w", err) } return nil }, ) if err != nil { return "", err } return token, nil } func (s AccountService) GetIdentityForMembership(ctx context.Context, membershipID gid.GID) (*coredata.Identity, error) { var ( scope = coredata.NewScopeFromObjectID(membershipID) identity = &coredata.Identity{} ) err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { membership := &coredata.Membership{} err := membership.LoadByID(ctx, conn, scope, membershipID) if err != nil { if err == coredata.ErrResourceNotFound { return NewMembershipNotFoundError(membershipID) } return fmt.Errorf("cannot load membership: %w", err) } err = identity.LoadByID(ctx, conn, membership.IdentityID) if err != nil { if err == coredata.ErrResourceNotFound { return NewIdentityNotFoundError(membership.IdentityID) } return fmt.Errorf("cannot load identity: %w", err) } return nil }, ) if err != nil { return nil, err } return identity, nil } func (s *AccountService) CreatePersonalAPIKey( ctx context.Context, identityID gid.GID, name string, expiresAt time.Time, ) (*coredata.PersonalAPIKey, string, error) { var ( personalAPIKey *coredata.PersonalAPIKey token string ) err := s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) (err error) { now := time.Now() personalAPIKey = &coredata.PersonalAPIKey{ ID: gid.New(gid.NilTenant, coredata.PersonalAPIKeyEntityType), IdentityID: identityID, Name: name, ExpiresAt: expiresAt, CreatedAt: now, UpdatedAt: now, } if err := personalAPIKey.Insert(ctx, tx); err != nil { return fmt.Errorf("cannot insert personal api key: %w", err) } token, err = securetoken.Sign( personalAPIKey.ID.String(), s.tokenSecret, ) if err != nil { return fmt.Errorf("cannot generate personal api key token: %w", err) } return nil }, ) if err != nil { return nil, "", err } return personalAPIKey, token, nil } func (s *AccountService) DeletePersonalAPIKey( ctx context.Context, identityID gid.GID, personalAPIKeyID gid.GID, ) error { return s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { personalAPIKey := &coredata.PersonalAPIKey{} err := personalAPIKey.LoadByID(ctx, tx, personalAPIKeyID) if err != nil { if err == coredata.ErrResourceNotFound { return NewPersonalAPIKeyNotFoundError(personalAPIKeyID) } return fmt.Errorf("cannot load personal api key: %w", err) } if personalAPIKey.IdentityID != identityID { return NewPersonalAPIKeyNotFoundError(personalAPIKeyID) } err = personalAPIKey.Delete(ctx, tx) if err != nil { return fmt.Errorf("cannot delete personal api key: %w", err) } return nil }, ) } func (s AccountService) ListInvitingOrganizations(ctx context.Context, identityID gid.GID) ([]*coredata.Organization, error) { var organizations coredata.Organizations err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { loaded, err := page.LoadAll( ctx, page.OrderBy[coredata.OrganizationOrderField]{ Field: coredata.OrganizationOrderFieldName, Direction: page.OrderDirectionAsc, }, func(ctx context.Context, cursor *page.Cursor[coredata.OrganizationOrderField]) ([]*coredata.Organization, error) { var batch coredata.Organizations if err := batch.LoadByIdentityIDWithPendingInvitation(ctx, conn, coredata.NewNoScope(), identityID, cursor); err != nil { return nil, fmt.Errorf("cannot load inviting organizations: %w", err) } return batch, nil }, ) if err != nil { return err } organizations = loaded return nil }, ) if err != nil { return nil, err } return organizations, nil } func (s AccountService) ListOrganizations(ctx context.Context, identityID gid.GID) ([]*coredata.Organization, error) { var organizations coredata.Organizations orderBy := page.OrderBy[coredata.OrganizationOrderField]{ Field: coredata.OrganizationOrderFieldCreatedAt, Direction: page.OrderDirectionDesc, } cursor := page.NewCursor(1000, nil, page.Head, orderBy) err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { err := organizations.LoadByIdentityID(ctx, conn, coredata.NewNoScope(), identityID, cursor) if err != nil { return fmt.Errorf("cannot load organizations: %w", err) } return nil }, ) if err != nil { return nil, err } return organizations, nil } func (s AccountService) GetMembershipForOrganization( ctx context.Context, identityID gid.GID, organizationID gid.GID, ) (*coredata.Membership, error) { membership := &coredata.Membership{} err := s.pg.WithTx( ctx, func(ctx context.Context, tx pg.Tx) error { identity := &coredata.Identity{} if err := identity.LoadByID(ctx, tx, identityID); err != nil { if errors.Is(err, coredata.ErrResourceNotFound) { return NewIdentityNotFoundError(identityID) } return fmt.Errorf("cannot load identity %q: %w", identityID, err) } if err := membership.LoadByIdentityIDAndOrganizationID( ctx, tx, coredata.NewScopeFromObjectID(organizationID), identityID, organizationID, ); err != nil { return fmt.Errorf("cannot load membership: %w", err) } return nil }, ) if err != nil { return nil, err } return membership, nil } func (s AccountService) ListSAMLConfigurationsForEmail( ctx context.Context, email mail.Addr, ) (coredata.SAMLConfigurations, error) { samlConfigurations := coredata.SAMLConfigurations{} err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { err := samlConfigurations.LoadVerifiedByEmailDomain(ctx, conn, email.Domain()) if err != nil { return fmt.Errorf("cannot load saml configurations: %w", err) } return nil }, ) if err != nil { return nil, err } return samlConfigurations, nil } func (s AccountService) CountSAMLConfigurationsForEmail( ctx context.Context, email mail.Addr, ) (int, error) { var ( count int samlConfigurations coredata.SAMLConfigurations ) err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) (err error) { count, err = samlConfigurations.CountVerifiedByEmailDomain(ctx, conn, email.Domain()) if err != nil { return fmt.Errorf("cannot count saml configurations: %w", err) } return nil }, ) if err != nil { return 0, err } return count, nil } func (s *AccountService) ListProfilesForIdentity( ctx context.Context, identityID gid.GID, cursor *page.Cursor[coredata.MembershipProfileOrderField], filter *coredata.MembershipProfileFilter, ) (*page.Page[*coredata.MembershipProfile, coredata.MembershipProfileOrderField], error) { var ( profiles = coredata.MembershipProfiles{} ) err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) error { if err := profiles.LoadByIdentityID(ctx, conn, identityID, cursor, filter); err != nil { return fmt.Errorf("cannot load profiles: %w", err) } return nil }, ) if err != nil { return nil, err } return page.NewPage(profiles, cursor), nil } func (s AccountService) CountProfiles( ctx context.Context, identityID gid.GID, filter *coredata.MembershipProfileFilter, ) (int, error) { var ( count int ) err := s.pg.WithConn( ctx, func(ctx context.Context, conn pg.Querier) (err error) { profiles := coredata.MembershipProfiles{} count, err = profiles.CountByIdentityID(ctx, conn, identityID, filter) if err != nil { return fmt.Errorf("cannot count profiles: %w", err) } return nil }, ) return count, err }