Upgrade to kit v0.3.0

Signed-off-by: Bryan Frimin <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2026-04-03 10:56:06 +02:00
parent 8adf26ad20
commit f17fb7bf49
191 changed files with 1617 additions and 1617 deletions

View File

@@ -99,7 +99,7 @@ func (s AccountService) ChangeEmail(ctx context.Context, identityID gid.GID, req
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
err := identity.LoadByID(ctx, tx, identityID)
if err != nil {
@@ -162,7 +162,7 @@ func (s AccountService) VerifyEmail(ctx context.Context, token string) error {
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
err := identity.LoadByID(ctx, tx, payload.Data.IdentityID)
if err != nil {
@@ -206,7 +206,7 @@ func (s *AccountService) ListPendingInvitations(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
profile := coredata.MembershipProfile{}
err := profile.LoadByID(ctx, conn, scope, userID)
if err != nil {
@@ -242,7 +242,7 @@ func (s AccountService) ChangePassword(ctx context.Context, identityID gid.GID,
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
err := identity.LoadByID(ctx, tx, identityID)
if err != nil {
@@ -287,7 +287,7 @@ func (s AccountService) CountSessions(ctx context.Context, identityID gid.GID) (
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
sessions := coredata.Sessions{}
count, err = sessions.CountByIdentityID(ctx, conn, identityID)
if err != nil {
@@ -310,7 +310,7 @@ func (s AccountService) ListSessions(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
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)
@@ -332,7 +332,7 @@ func (s AccountService) GetIdentity(ctx context.Context, identityID gid.GID) (*c
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := identity.LoadByID(ctx, conn, identityID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -361,7 +361,7 @@ func (s AccountService) UpdateIdentity(ctx context.Context, identityID gid.GID,
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := identity.LoadByID(ctx, tx, identityID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -397,7 +397,7 @@ func (s AccountService) ListPersonalAPIKeys(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
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)
@@ -419,7 +419,7 @@ func (s AccountService) CountPersonalAPIKeys(ctx context.Context, identityID gid
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
personalAccessTokens := coredata.PersonalAPIKeys{}
count, err = personalAccessTokens.CountByIdentityID(ctx, conn, identityID)
if err != nil {
@@ -441,7 +441,7 @@ func (s *AccountService) RevealPersonalAPIKeyToken(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) (err error) {
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 {
@@ -481,7 +481,7 @@ func (s AccountService) GetIdentityForMembership(ctx context.Context, membership
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
membership := &coredata.Membership{}
err := membership.LoadByID(ctx, conn, scope, membershipID)
if err != nil {
@@ -525,7 +525,7 @@ func (s *AccountService) CreatePersonalAPIKey(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) (err error) {
func(ctx context.Context, tx pg.Tx) (err error) {
now := time.Now()
personalAPIKey = &coredata.PersonalAPIKey{
@@ -567,7 +567,7 @@ func (s *AccountService) DeletePersonalAPIKey(
) error {
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
personalAPIKey := &coredata.PersonalAPIKey{}
err := personalAPIKey.LoadByID(ctx, tx, personalAPIKeyID)
if err != nil {
@@ -602,7 +602,7 @@ func (s AccountService) ListOrganizations(ctx context.Context, identityID gid.GI
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
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)
@@ -628,7 +628,7 @@ func (s AccountService) GetMembershipForOrganization(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
if err := identity.LoadByID(ctx, tx, identityID); err != nil {
@@ -668,7 +668,7 @@ func (s AccountService) ListSAMLConfigurationsForEmail(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
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)
@@ -695,7 +695,7 @@ func (s AccountService) CountSAMLConfigurationsForEmail(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
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)
@@ -723,7 +723,7 @@ func (s *AccountService) ListProfilesForIdentity(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
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)
}
@@ -750,7 +750,7 @@ func (s AccountService) CountProfiles(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
profiles := coredata.MembershipProfiles{}
count, err = profiles.CountByIdentityID(ctx, conn, identityID, filter)
if err != nil {

View File

@@ -42,7 +42,7 @@ func (s *APIKeyService) GetAPIKey(ctx context.Context, keyID gid.GID) (*coredata
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
if err := apiKey.LoadByID(ctx, tx, keyID); err != nil {
if err == coredata.ErrResourceNotFound {
return NewPersonalAPIKeyNotFoundError(keyID)

View File

@@ -142,7 +142,7 @@ func (s *AuthService) ActivateAccount(
if err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := invitation.LoadByID(ctx, tx, scope, payload.Data.InvitationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -256,7 +256,7 @@ func (s AuthService) ResetPassword(
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
err := identity.LoadByEmail(ctx, tx, payload.Data.Email)
if err != nil {
@@ -295,7 +295,7 @@ func (s AuthService) SendPasswordResetInstructionByEmail(
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
if err := identity.LoadByEmail(ctx, tx, email); err != nil {
if err == coredata.ErrResourceNotFound {
@@ -396,7 +396,7 @@ func (s AuthService) CreateIdentityWithPassword(
err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := identity.Insert(ctx, tx)
if err != nil {
if err == coredata.ErrResourceAlreadyExists {
@@ -426,7 +426,7 @@ func (s AuthService) OpenSessionWithSAML(ctx context.Context, identityID gid.GID
err := s.pg.WithTx(
ctx,
func(conn pg.Conn) (err error) {
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 {
@@ -449,7 +449,7 @@ func (s AuthService) OpenSessionWithOIDC(ctx context.Context, identityID gid.GID
err := s.pg.WithTx(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Tx) (err error) {
session = coredata.NewRootSession(identityID, authMethod, s.sessionDuration)
err = session.Insert(ctx, conn)
if err != nil {
@@ -484,7 +484,7 @@ func (s AuthService) CheckCredentials(
err = s.pg.WithTx(
ctx,
func(conn pg.Conn) error {
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
@@ -521,7 +521,7 @@ func (s AuthService) OpenSessionWithPassword(ctx context.Context, identityID gid
err := s.pg.WithTx(
ctx,
func(conn pg.Conn) (err error) {
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 {
@@ -555,7 +555,7 @@ func (s AuthService) SendMagicLink(ctx context.Context, req *SendMagicLinkReques
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
hashedToken := HashToken(tokenString)
token := &coredata.Token{
ID: gid.New(gid.NilTenant, coredata.TokenEntityType),
@@ -664,7 +664,7 @@ func (s AuthService) OpenSessionWithMagicLink(ctx context.Context, tokenString s
if err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
hashedValue := HashToken(tokenString)
token := &coredata.Token{}

View File

@@ -33,7 +33,7 @@ import (
// AuthorizationAttributer is implemented by entities that provide attributes
// for policy condition evaluation.
type AuthorizationAttributer interface {
AuthorizationAttributes(ctx context.Context, conn pg.Conn) (map[string]string, error)
AuthorizationAttributes(ctx context.Context, conn pg.Querier) (map[string]string, error)
}
// AuthorizeParams contains the parameters for an authorization request.
@@ -76,11 +76,11 @@ func (a *Authorizer) Authorize(ctx context.Context, params AuthorizeParams) erro
return NewUnsupportedPrincipalTypeError(params.Principal.EntityType())
}
return a.pg.WithConn(ctx, func(conn pg.Conn) error { return a.authorize(ctx, conn, params) })
return a.pg.WithTx(ctx, func(ctx context.Context, tx pg.Tx) error { return a.authorize(ctx, tx, params) })
}
func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params AuthorizeParams) error {
resourceAttrs, err := a.buildResourceAttributes(ctx, conn, params)
func (a *Authorizer) authorize(ctx context.Context, tx pg.Tx, params AuthorizeParams) error {
resourceAttrs, err := a.buildResourceAttributes(ctx, tx, params)
if err != nil {
return fmt.Errorf("cannot build resource attributes: %w", err)
}
@@ -88,7 +88,7 @@ func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params Authori
resourceOrgID := resourceAttrs["organization_id"]
// Find role for resource's organization
membership, err := a.loadMembership(ctx, conn, params.Principal, resourceOrgID)
membership, err := a.loadMembership(ctx, tx, params.Principal, resourceOrgID)
if err != nil {
return fmt.Errorf("cannot load memberships for principal: %w", err)
}
@@ -97,7 +97,7 @@ func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params Authori
if membership != nil && params.Session != nil && !params.SkipAssumptionCheck {
if _, err := a.getActiveChildSessionForMembership(
ctx,
conn,
tx,
*params.Session,
membership.ID,
); err != nil {
@@ -126,7 +126,7 @@ func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params Authori
}
}
principalAttrs, err := a.buildPrincipalAttributes(ctx, conn, params.Principal, scopedPrincipalAttrs)
principalAttrs, err := a.buildPrincipalAttributes(ctx, tx, params.Principal, scopedPrincipalAttrs)
if err != nil {
return fmt.Errorf("cannot build principal attributes: %w", err)
}
@@ -144,7 +144,7 @@ func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params Authori
}
if a.evaluator.Evaluate(req, policies).IsAllowed() {
a.recordAuditLog(ctx, conn, params, resourceAttrs)
a.recordAuditLog(ctx, tx, params, resourceAttrs)
return nil
}
@@ -153,7 +153,7 @@ func (a *Authorizer) authorize(ctx context.Context, conn pg.Conn, params Authori
func (a *Authorizer) loadMembership(
ctx context.Context,
conn pg.Conn,
conn pg.Querier,
principalID gid.GID,
resourceOrgID string,
) (*coredata.Membership, error) {
@@ -180,7 +180,7 @@ func (a *Authorizer) loadMembership(
func (a *Authorizer) getActiveChildSessionForMembership(
ctx context.Context,
conn pg.Conn,
conn pg.Querier,
rootSessionID gid.GID,
membershipID gid.GID,
) (*coredata.Session, error) {
@@ -203,7 +203,7 @@ func (a *Authorizer) getActiveChildSessionForMembership(
func (a *Authorizer) buildPrincipalAttributes(
ctx context.Context,
conn pg.Conn,
conn pg.Querier,
principalID gid.GID,
defaultAttrs map[string]string,
) (map[string]string, error) {
@@ -227,7 +227,7 @@ func (a *Authorizer) buildPrincipalAttributes(
func (a *Authorizer) buildResourceAttributes(
ctx context.Context,
conn pg.Conn,
conn pg.Querier,
params AuthorizeParams,
) (map[string]string, error) {
attrs := map[string]string{
@@ -288,7 +288,7 @@ func resourceTypeFromAction(action string) string {
func (a *Authorizer) recordAuditLog(
ctx context.Context,
conn pg.Conn,
tx pg.Tx,
params AuthorizeParams,
resourceAttrs map[string]string,
) {
@@ -344,7 +344,7 @@ func (a *Authorizer) recordAuditLog(
scope := coredata.NewScope(orgID.TenantID())
if err := entry.Insert(ctx, conn, scope); err != nil {
if err := entry.Insert(ctx, tx, scope); err != nil {
a.logger.ErrorCtx(
ctx,
"cannot insert audit log entry",

View File

@@ -49,7 +49,7 @@ func (s *CompliancePageService) GenerateLogoURL(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
if err := compliancePage.LoadByID(ctx, conn, scope, compliancePageID); err != nil {
return fmt.Errorf("cannot load compliance page: %w", err)
}
@@ -98,7 +98,7 @@ func (s *CompliancePageService) EmailPresenterConfig(ctx context.Context, compli
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
if err := compliancePage.LoadByID(ctx, conn, scope, compliancePageID); err != nil {
return fmt.Errorf("cannot load compliance page: %w", err)
}

View File

@@ -92,7 +92,7 @@ func (gc *GarbageCollector) cleanup(ctx context.Context) error {
return gc.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
var state coredata.OIDCState
deleted, err := state.DeleteExpired(ctx, tx, now)
if err != nil {

View File

@@ -302,7 +302,7 @@ func (s *Service) InitiateLogin(
err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
if err := oidcState.Insert(ctx, tx); err != nil {
return fmt.Errorf("cannot store oidc state: %w", err)
}
@@ -339,7 +339,7 @@ func (s *Service) HandleCallback(
var oidcState coredata.OIDCState
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
if err := oidcState.LoadByIDForUpdate(ctx, tx, stateParam); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return NewInvalidStateError()
@@ -407,7 +407,7 @@ func (s *Service) HandleCallback(
err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity = &coredata.Identity{}
err := identity.LoadByEmail(ctx, tx, email)
if err != nil {

View File

@@ -252,7 +252,7 @@ func (s *OrganizationService) UpdateMempership(
membership := coredata.Membership{}
if err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
if err := membership.LoadByID(ctx, tx, scope, membershipID); err != nil {
if err == coredata.ErrResourceNotFound {
@@ -291,7 +291,7 @@ func (s *OrganizationService) RemoveUser(
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
profile := coredata.MembershipProfile{}
if err := profile.LoadByID(ctx, tx, scope, profileID); err != nil {
@@ -359,7 +359,7 @@ func (s *OrganizationService) InviteUser(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
organization := coredata.Organization{}
err := organization.LoadByID(ctx, tx, scope, req.OrganizationID)
if err != nil {
@@ -583,7 +583,7 @@ func (s *OrganizationService) CreateOrganization(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
err := identity.LoadByID(ctx, tx, identityID)
if err != nil {
@@ -764,7 +764,7 @@ func (s *OrganizationService) UpdateOrganization(ctx context.Context, organizati
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := organization.LoadByID(ctx, tx, scope, organizationID)
if err != nil {
return fmt.Errorf("cannot load organization: %w", err)
@@ -852,7 +852,7 @@ func (s *OrganizationService) DeleteOrganization(ctx context.Context, organizati
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
organization := &coredata.Organization{}
err := organization.LoadByID(ctx, tx, scope, organizationID)
if err != nil {
@@ -882,7 +882,7 @@ func (s *OrganizationService) CreateUser(ctx context.Context, req *CreateUserReq
err := s.pg.WithTx(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Tx) error {
identity := &coredata.Identity{}
if err := identity.LoadByEmail(ctx, conn, req.EmailAddress); err != nil {
if !errors.Is(err, coredata.ErrResourceNotFound) {
@@ -973,7 +973,7 @@ func (s *OrganizationService) UpdateUser(ctx context.Context, req *UpdateUserReq
err := s.pg.WithTx(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Tx) error {
if err := profile.LoadByID(ctx, conn, scope, req.ID); err != nil {
return fmt.Errorf("cannot load profile: %w", err)
}
@@ -1024,17 +1024,17 @@ func (s *OrganizationService) UpdateUserState(
profile = &coredata.MembershipProfile{}
)
err := s.pg.WithConn(
err := s.pg.WithTx(
ctx,
func(conn pg.Conn) error {
if err := profile.LoadByID(ctx, conn, scope, userID); err != nil {
func(ctx context.Context, tx pg.Tx) error {
if err := profile.LoadByID(ctx, tx, scope, userID); err != nil {
return fmt.Errorf("cannot load profile: %w", err)
}
profile.State = state
profile.UpdatedAt = time.Now()
if err := profile.Update(ctx, conn, scope); err != nil {
if err := profile.Update(ctx, tx, scope); err != nil {
return fmt.Errorf("cannot update profile: %w", err)
}
@@ -1054,7 +1054,7 @@ func (s *OrganizationService) GetProfile(ctx context.Context, profileID gid.GID)
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
if err := profile.LoadByID(ctx, conn, coredata.NewScopeFromObjectID(profileID), profileID); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return NewProfileNotFoundError(profileID)
@@ -1083,7 +1083,7 @@ func (s *OrganizationService) GetProfilesByIDs(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
if err := profiles.LoadByIDs(
ctx,
conn,
@@ -1108,7 +1108,7 @@ func (s *OrganizationService) GetProfileForIdentityAndOrganization(ctx context.C
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
if err := profile.LoadByIdentityIDAndOrganizationID(
ctx,
conn,
@@ -1147,7 +1147,7 @@ func (s *OrganizationService) ListProfiles(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
if err := profiles.LoadByOrganizationID(ctx, conn, scope, organizationID, cursor, filter); err != nil {
return fmt.Errorf("cannot load profiles: %w", err)
}
@@ -1175,7 +1175,7 @@ func (s OrganizationService) CountProfiles(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
profiles := coredata.MembershipProfiles{}
count, err = profiles.CountByOrganizationID(ctx, conn, scope, organizationID, filter)
if err != nil {
@@ -1197,7 +1197,7 @@ func (s *OrganizationService) GetOrganizationForMembership(ctx context.Context,
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
membership := &coredata.Membership{}
err := membership.LoadByID(ctx, conn, scope, membershipID)
if err != nil {
@@ -1241,7 +1241,7 @@ func (s OrganizationService) GenerateLogoURL(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
organization := &coredata.Organization{}
if err := organization.LoadByID(ctx, conn, scope, organizationID); err != nil {
return fmt.Errorf("cannot load organization: %w", err)
@@ -1287,7 +1287,7 @@ func (s OrganizationService) GenerateHorizontalLogoURL(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
organization := &coredata.Organization{}
if err := organization.LoadByID(ctx, conn, scope, organizationID); err != nil {
return fmt.Errorf("cannot load organization: %w", err)
@@ -1329,7 +1329,7 @@ func (s OrganizationService) DeleteSAMLConfiguration(
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
var config coredata.SAMLConfiguration
if err := config.LoadByID(ctx, tx, scope, configID); err != nil {
return fmt.Errorf("cannot load saml configuration: %w", err)
@@ -1360,7 +1360,7 @@ func (s OrganizationService) ListSAMLConfigurations(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := samlConfigurations.LoadByOrganizationID(ctx, conn, scope, organizationID)
if err != nil {
return fmt.Errorf("cannot load saml configurations: %w", err)
@@ -1387,7 +1387,7 @@ func (s OrganizationService) CountSAMLConfigurations(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
samlConfigurations := coredata.SAMLConfigurations{}
count, err = samlConfigurations.CountByOrganizationID(ctx, conn, scope, organizationID)
if err != nil {
@@ -1413,7 +1413,7 @@ func (s OrganizationService) ListSCIMEvents(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := scimEvents.LoadByOrganizationID(ctx, conn, scope, organizationID, cursor)
if err != nil {
return fmt.Errorf("cannot load scim events: %w", err)
@@ -1440,7 +1440,7 @@ func (s OrganizationService) CountSCIMEvents(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
scimEvents := coredata.SCIMEvents{}
count, err = scimEvents.CountByOrganizationID(ctx, conn, scope, organizationID)
if err != nil {
@@ -1465,7 +1465,7 @@ func (s OrganizationService) GetSCIMConfiguration(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := config.LoadByOrganizationID(ctx, conn, scope, organizationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -1509,7 +1509,7 @@ func (s OrganizationService) CreateSCIMConfiguration(
err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := config.Insert(ctx, tx, scope)
if err != nil {
if err == coredata.ErrResourceAlreadyExists {
@@ -1537,7 +1537,7 @@ func (s OrganizationService) DeleteSCIMConfiguration(
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
config := &coredata.SCIMConfiguration{}
err := config.LoadByID(ctx, tx, scope, configID)
if err != nil {
@@ -1608,7 +1608,7 @@ func (s OrganizationService) RegenerateSCIMToken(
err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := config.LoadByID(ctx, tx, scope, configID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -1652,7 +1652,7 @@ func (s OrganizationService) UpdateSCIMBridge(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := bridge.LoadByID(ctx, tx, scope, bridgeID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -1697,7 +1697,7 @@ func (s OrganizationService) ListSCIMEventsByConfigID(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := scimEvents.LoadBySCIMConfigurationID(ctx, conn, scope, scimConfigurationID, cursor)
if err != nil {
return fmt.Errorf("cannot load scim events: %w", err)
@@ -1724,7 +1724,7 @@ func (s OrganizationService) CountSCIMEventsByConfigID(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
scimEvents := coredata.SCIMEvents{}
count, err = scimEvents.CountBySCIMConfigurationID(ctx, conn, scope, scimConfigurationID)
if err != nil {
@@ -1784,7 +1784,7 @@ func (s OrganizationService) CreateSAMLConfiguration(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
organization := &coredata.Organization{}
err := organization.LoadByID(ctx, tx, scope, organizationID)
if err != nil {
@@ -1823,7 +1823,7 @@ func (s OrganizationService) UpdateSAMLConfiguration(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
organization := &coredata.Organization{}
err := organization.LoadByID(ctx, tx, scope, organizationID)
if err != nil {
@@ -1902,7 +1902,7 @@ func (s OrganizationService) GetOrganization(ctx context.Context, organizationID
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := organization.LoadByID(ctx, conn, scope, organizationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -1930,7 +1930,7 @@ func (s OrganizationService) GetSCIMBridgeByID(ctx context.Context, bridgeID gid
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := bridge.LoadByID(ctx, conn, scope, bridgeID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -1961,7 +1961,7 @@ func (s OrganizationService) GetConnectorMetadataByID(ctx context.Context, conne
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := connector.LoadMetadataByID(ctx, conn, scope, connectorID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -1990,7 +1990,7 @@ func (s OrganizationService) GetSCIMBridgeByOrganizationID(ctx context.Context,
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := bridge.LoadByOrganizationID(ctx, conn, scope, organizationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -2030,7 +2030,7 @@ func (s OrganizationService) CreateSCIMBridge(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
organization := &coredata.Organization{}
err := organization.LoadByID(ctx, tx, scope, organizationID)
if err != nil {
@@ -2113,7 +2113,7 @@ func (s OrganizationService) DeleteSCIMBridge(ctx context.Context, organizationI
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
organization := &coredata.Organization{}
err := organization.LoadByID(ctx, tx, scope, organizationID)
if err != nil {
@@ -2154,7 +2154,7 @@ func (s *OrganizationService) GetAuditLogEntry(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
return entry.LoadByID(ctx, conn, scope, id)
},
)
@@ -2178,7 +2178,7 @@ func (s *OrganizationService) ListAuditLogEntries(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
if err := entries.LoadAllByOrganizationID(ctx, conn, scope, organizationID, cursor, filter); err != nil {
return fmt.Errorf("cannot load audit log entries: %w", err)
}
@@ -2205,7 +2205,7 @@ func (s *OrganizationService) CountAuditLogEntries(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
func(ctx context.Context, conn pg.Querier) (err error) {
entries := coredata.AuditLogEntries{}
count, err = entries.CountByOrganizationID(ctx, conn, scope, organizationID, filter)
if err != nil {

View File

@@ -96,7 +96,7 @@ func (gc *GarbageCollector) cleanup(ctx context.Context) error {
return gc.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
assertionsDeleted, err := coredata.DeleteExpiredSAMLAssertions(ctx, tx, now)
if err != nil {
return fmt.Errorf("cannot delete expired saml assertions: %w", err)

View File

@@ -110,7 +110,7 @@ func (s *Service) InitiateLogin(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
config := &coredata.SAMLConfiguration{}
err := config.LoadByID(ctx, tx, coredata.NewNoScope(), configID)
if err != nil {
@@ -177,7 +177,7 @@ func (s *Service) HandleAssertion(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
config := &coredata.SAMLConfiguration{}
err := config.LoadByID(ctx, tx, coredata.NewNoScope(), configID)

View File

@@ -101,7 +101,7 @@ func (v *SAMLDomainVerifier) checkUnverifiedDomains(ctx context.Context) error {
err := v.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := configs.LoadUnverified(ctx, conn)
if err != nil {
return fmt.Errorf("cannot load unverified SAML configurations: %w", err)
@@ -148,7 +148,7 @@ func (v *SAMLDomainVerifier) checkUnverifiedDomains(ctx context.Context) error {
func (v *SAMLDomainVerifier) tryVerifyDomain(ctx context.Context, configID gid.GID) error {
return v.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
config := &coredata.SAMLConfiguration{}
if err := config.LoadByIDForUpdateSkipLocked(ctx, tx, configID); err != nil {
if err == coredata.ErrResourceNotFound {

View File

@@ -38,7 +38,7 @@ func (r *BridgeRunner) acquireNextBridge(ctx context.Context) (*coredata.SCIMBri
err := r.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
bridge = &coredata.SCIMBridge{}
if err := bridge.LoadNextForSyncSkipLocked(ctx, tx, r.cfg.StaleSyncThreshold); err != nil {
return err
@@ -70,9 +70,9 @@ func (r *BridgeRunner) transitionToSuccess(
connector *coredata.Connector,
logger *log.Logger,
) error {
return r.pg.WithConn(
return r.pg.WithTx(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
now := time.Now()
nextSync := now.Add(r.cfg.Interval)
@@ -84,7 +84,7 @@ func (r *BridgeRunner) transitionToSuccess(
bridge.TotalSyncCount++
bridge.UpdatedAt = now
if err := bridge.Update(ctx, conn, scope); err != nil {
if err := bridge.Update(ctx, tx, scope); err != nil {
logger.ErrorCtx(
ctx,
"cannot update bridge after successful sync",
@@ -95,7 +95,7 @@ func (r *BridgeRunner) transitionToSuccess(
if connector != nil {
connector.UpdatedAt = now
if err := connector.Update(ctx, conn, scope, r.encryptionKey); err != nil {
if err := connector.Update(ctx, tx, scope, r.encryptionKey); err != nil {
logger.WarnCtx(
ctx,
"cannot persist refreshed OAuth2 token",
@@ -129,9 +129,9 @@ func (r *BridgeRunner) transitionToFailed(
duration time.Duration,
logger *log.Logger,
) error {
return r.pg.WithConn(
return r.pg.WithTx(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
now := time.Now()
bridge.ConsecutiveFailures++
@@ -171,7 +171,7 @@ func (r *BridgeRunner) transitionToFailed(
)
}
if err := bridge.Update(ctx, conn, scope); err != nil {
if err := bridge.Update(ctx, tx, scope); err != nil {
logger.ErrorCtx(
ctx,
"cannot update bridge after failed sync",

View File

@@ -38,11 +38,11 @@ func (r *BridgeRunner) executeSync(
) (stats SyncStats, duration time.Duration, connector *coredata.Connector, err error) {
start := time.Now()
err = r.pg.WithConn(
err = r.pg.WithTx(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
var syncErr error
stats, connector, syncErr = r.doSync(ctx, conn, bridge, scope, logger)
stats, connector, syncErr = r.doSync(ctx, tx, bridge, scope, logger)
return syncErr
},
)
@@ -53,7 +53,7 @@ func (r *BridgeRunner) executeSync(
func (r *BridgeRunner) doSync(
ctx context.Context,
conn pg.Conn,
tx pg.Tx,
scimBridge *coredata.SCIMBridge,
scope coredata.Scoper,
logger *log.Logger,
@@ -63,7 +63,7 @@ func (r *BridgeRunner) doSync(
}
dbConnector := &coredata.Connector{}
if err := dbConnector.LoadByID(ctx, conn, scope, *scimBridge.ConnectorID, r.encryptionKey); err != nil {
if err := dbConnector.LoadByID(ctx, tx, scope, *scimBridge.ConnectorID, r.encryptionKey); err != nil {
return SyncStats{}, nil, fmt.Errorf("cannot load connector: %w", err)
}
@@ -73,7 +73,7 @@ func (r *BridgeRunner) doSync(
}
var scimConfig coredata.SCIMConfiguration
if err := scimConfig.LoadByID(ctx, conn, scope, scimBridge.ScimConfigurationID); err != nil {
if err := scimConfig.LoadByID(ctx, tx, scope, scimBridge.ScimConfigurationID); err != nil {
return SyncStats{}, nil, fmt.Errorf("cannot load SCIM configuration: %w", err)
}
@@ -84,7 +84,7 @@ func (r *BridgeRunner) doSync(
scimConfig.HashedToken = HashToken(token)
scimConfig.UpdatedAt = time.Now()
if err := scimConfig.Update(ctx, conn, scope); err != nil {
if err := scimConfig.Update(ctx, tx, scope); err != nil {
return SyncStats{}, nil, fmt.Errorf("cannot update SCIM configuration token: %w", err)
}

View File

@@ -107,7 +107,7 @@ func (s *Service) ValidateToken(ctx context.Context, token string) (*coredata.SC
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := config.LoadByHashedToken(ctx, conn, hashedToken)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -160,7 +160,7 @@ func (s *Service) CreateUser(
scope := coredata.NewScopeFromObjectID(config.OrganizationID)
err = s.pg.WithTx(ctx, func(tx pg.Conn) error {
err = s.pg.WithTx(ctx, func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
if err := identity.LoadByEmail(ctx, tx, emailAddr); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
@@ -355,7 +355,7 @@ func (s *Service) GetUser(
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
profile = &coredata.MembershipProfile{}
if err := profile.LoadByID(ctx, conn, scope, profileID); err != nil {
if err == coredata.ErrResourceNotFound {
@@ -405,7 +405,7 @@ func (s *Service) ListUsers(
err = s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
var err error
totalCount, err = profiles.CountByOrganizationID(ctx, conn, scope, config.OrganizationID, filter)
if err != nil {
@@ -486,7 +486,7 @@ func (s *Service) updateUser(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
profile = &coredata.MembershipProfile{}
if err := profile.LoadByID(ctx, tx, scope, profileID); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
@@ -772,7 +772,7 @@ func (s *Service) DeleteUser(
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
profile := &coredata.MembershipProfile{}
if err := profile.LoadByID(ctx, tx, scope, profileID); err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
@@ -837,10 +837,10 @@ func (s *Service) LogEvent(
event := s.createEvent(config, method, path, userName, ipAddress, statusCode, errorMessage)
err := s.pg.WithConn(
err := s.pg.WithTx(
ctx,
func(conn pg.Conn) error {
err := event.Insert(ctx, conn, scope)
func(ctx context.Context, tx pg.Tx) error {
err := event.Insert(ctx, tx, scope)
if err != nil {
return fmt.Errorf("cannot insert SCIM event: %w", err)
}

View File

@@ -249,7 +249,7 @@ func (s *Service) GetMembership(ctx context.Context, membershipID gid.GID) (*cor
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := membership.LoadByID(ctx, conn, scope, membershipID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -277,7 +277,7 @@ func (s *Service) GetInvitation(ctx context.Context, invitationID gid.GID) (*cor
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := invitation.LoadByID(ctx, conn, scope, invitationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -302,7 +302,7 @@ func (s *Service) GetSession(ctx context.Context, sessionID gid.GID) (*coredata.
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := session.LoadByID(ctx, conn, sessionID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -330,7 +330,7 @@ func (s *Service) GetSAMLconfiguration(ctx context.Context, samlConfigurationID
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := samlConfiguration.LoadByID(ctx, conn, scope, samlConfigurationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -355,7 +355,7 @@ func (s *Service) GetPersonalAPIKey(ctx context.Context, personalAPIKeyID gid.GI
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := personalAPIKey.LoadByID(ctx, conn, personalAPIKeyID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -383,7 +383,7 @@ func (s *Service) GetSCIMConfiguration(ctx context.Context, scimConfigurationID
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := scimConfiguration.LoadByID(ctx, conn, scope, scimConfigurationID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -411,7 +411,7 @@ func (s *Service) GetSCIMEvent(ctx context.Context, scimEventID gid.GID) (*cored
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Querier) error {
err := scimEvent.LoadByID(ctx, conn, scope, scimEventID)
if err != nil {
if err == coredata.ErrResourceNotFound {

View File

@@ -43,7 +43,7 @@ func (s SessionService) GetSession(ctx context.Context, sessionID gid.GID) (*cor
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
if err := session.LoadByID(ctx, tx, sessionID); err != nil {
if err == coredata.ErrResourceNotFound {
return NewSessionNotFoundError(sessionID)
@@ -79,7 +79,7 @@ func (s SessionService) GetSession(ctx context.Context, sessionID gid.GID) (*cor
func (s SessionService) CloseSession(ctx context.Context, sessionID gid.GID) error {
return s.pg.WithTx(
ctx,
func(conn pg.Conn) error {
func(ctx context.Context, conn pg.Tx) error {
session := &coredata.Session{}
if err := session.LoadByID(ctx, conn, sessionID); err != nil {
if err == coredata.ErrResourceNotFound {
@@ -114,7 +114,7 @@ func (s SessionService) RevokeSession(ctx context.Context, identityID gid.GID, s
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
identity := &coredata.Identity{}
err := identity.LoadByID(ctx, tx, identityID)
if err != nil {
@@ -166,7 +166,7 @@ func (s SessionService) RevokeAllSessions(ctx context.Context, currentSessionID
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
session := coredata.Session{}
err := session.LoadByID(ctx, tx, currentSessionID)
if err != nil {
@@ -193,7 +193,7 @@ func (s SessionService) RevokeAllSessions(ctx context.Context, currentSessionID
func (s SessionService) UpdateSessionInfo(ctx context.Context, sessionID gid.GID, userAgent string, ipAddress net.IP) error {
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
session := &coredata.Session{}
err := session.LoadByID(ctx, tx, sessionID)
if err != nil {
@@ -224,7 +224,7 @@ func (s SessionService) UpdateSessionInfo(ctx context.Context, sessionID gid.GID
func (s SessionService) UpdateSessionData(ctx context.Context, sessionID gid.GID, data coredata.SessionData) error {
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
session := &coredata.Session{}
err := session.LoadByID(ctx, tx, sessionID)
if err != nil {
@@ -256,7 +256,7 @@ func (s SessionService) GetActiveSessionForMembership(ctx context.Context, rootS
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
rootSession := &coredata.Session{}
err := rootSession.LoadByID(ctx, tx, rootSessionID)
if err != nil {
@@ -318,7 +318,7 @@ func (s SessionService) OpenPasswordChildSessionForOrganization(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := rootSession.LoadByID(ctx, tx, rootSessionID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -420,7 +420,7 @@ func (s SessionService) OpenSAMLChildSessionForOrganization(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
err := rootSession.LoadByID(ctx, tx, rootSessionID)
if err != nil {
if err == coredata.ErrResourceNotFound {
@@ -509,7 +509,7 @@ func (s SessionService) AssumeOrganizationSession(
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
func(ctx context.Context, tx pg.Tx) error {
if err := rootSession.LoadByID(ctx, tx, sessionID); err != nil {
if err == coredata.ErrResourceNotFound {
return NewSessionNotFoundError(sessionID)