Enforce allowed tenant in each request
Signed-off-by: gearnode <bryan@frimin.fr>
This commit is contained in:
@@ -264,44 +264,6 @@ func (s Service) GetSession(
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func (s Service) RefreshSession(
|
||||
ctx context.Context,
|
||||
sessionID gid.GID,
|
||||
) (*coredata.Session, error) {
|
||||
session := &coredata.Session{}
|
||||
|
||||
err := s.pg.WithTx(
|
||||
ctx,
|
||||
func(tx pg.Conn) error {
|
||||
if err := session.LoadByID(ctx, tx, sessionID); err != nil {
|
||||
return &ErrSessionNotFound{message: "session not found"}
|
||||
}
|
||||
|
||||
// Check if session is expired
|
||||
if time.Now().After(session.ExpiredAt) {
|
||||
return &ErrSessionExpired{message: "session expired"}
|
||||
}
|
||||
|
||||
// Update session expiration
|
||||
now := time.Now()
|
||||
session.ExpiredAt = now.Add(24 * time.Hour)
|
||||
session.UpdatedAt = now
|
||||
|
||||
if err := session.Update(ctx, tx); err != nil {
|
||||
return fmt.Errorf("cannot update session: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
},
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return session, nil
|
||||
}
|
||||
|
||||
func (s Service) GetUserByID(
|
||||
ctx context.Context,
|
||||
userID gid.GID,
|
||||
@@ -337,43 +299,28 @@ func (s Service) GetUserBySession(
|
||||
return s.GetUserByID(ctx, session.UserID)
|
||||
}
|
||||
|
||||
// GetUserOrganizations gets all organizations for a user
|
||||
func (s Service) GetUserOrganizations(
|
||||
func (s Service) ListOrganizationsForUserID(
|
||||
ctx context.Context,
|
||||
userID gid.GID,
|
||||
) ([]gid.GID, error) {
|
||||
q := `
|
||||
SELECT
|
||||
organization_id
|
||||
FROM
|
||||
usrmgr_user_organizations
|
||||
WHERE
|
||||
user_id = @user_id;
|
||||
`
|
||||
) (coredata.Organizations, error) {
|
||||
|
||||
args := pgx.StrictNamedArgs{"user_id": userID}
|
||||
uos := coredata.UserOrganizations{}
|
||||
organizations := []*coredata.Organization{}
|
||||
|
||||
var organizationIDs []gid.GID
|
||||
|
||||
err := s.pg.WithTx(
|
||||
err := s.pg.WithConn(
|
||||
ctx,
|
||||
func(tx pg.Conn) error {
|
||||
rows, err := tx.Query(ctx, q, args)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to query user organizations: %w", err)
|
||||
func(conn pg.Conn) error {
|
||||
if err := uos.ForUserID(ctx, conn, userID); err != nil {
|
||||
return fmt.Errorf("cannot list user organizations: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
var organizationID gid.GID
|
||||
if err := rows.Scan(&organizationID); err != nil {
|
||||
return fmt.Errorf("failed to scan organization ID: %w", err)
|
||||
for _, uo := range uos {
|
||||
scope := coredata.NewScope(uo.OrganizationID.TenantID())
|
||||
organization := &coredata.Organization{}
|
||||
if err := organization.LoadByID(ctx, conn, scope, uo.OrganizationID); err != nil {
|
||||
return fmt.Errorf("cannot load organization by id: %w", err)
|
||||
}
|
||||
organizationIDs = append(organizationIDs, organizationID)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return fmt.Errorf("error iterating over rows: %w", err)
|
||||
organizations = append(organizations, organization)
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -384,58 +331,55 @@ WHERE
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return organizationIDs, nil
|
||||
return organizations, nil
|
||||
}
|
||||
|
||||
// AddUserToOrganization adds a user to an organization
|
||||
func (s Service) AddUserToOrganization(
|
||||
func (s Service) ListTenantsForUserID(
|
||||
ctx context.Context,
|
||||
userID gid.GID,
|
||||
) ([]gid.TenantID, error) {
|
||||
|
||||
uos := coredata.UserOrganizations{}
|
||||
|
||||
err := s.pg.WithConn(
|
||||
ctx,
|
||||
func(tx pg.Conn) error {
|
||||
return uos.ForUserID(ctx, tx, userID)
|
||||
},
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
tenantIDs := make([]gid.TenantID, len(uos))
|
||||
for _, uo := range uos {
|
||||
tenantIDs = append(tenantIDs, uo.OrganizationID.TenantID())
|
||||
}
|
||||
|
||||
return tenantIDs, nil
|
||||
}
|
||||
|
||||
func (s Service) EnrollUserInOrganization(
|
||||
ctx context.Context,
|
||||
userID gid.GID,
|
||||
organizationID gid.GID,
|
||||
) error {
|
||||
q := `
|
||||
INSERT INTO
|
||||
usrmgr_user_organizations (user_id, organization_id, created_at)
|
||||
VALUES
|
||||
(@user_id, @organization_id, NOW())
|
||||
ON CONFLICT (user_id, organization_id) DO NOTHING;
|
||||
`
|
||||
|
||||
args := pgx.StrictNamedArgs{
|
||||
"user_id": userID,
|
||||
"organization_id": organizationID,
|
||||
uo := coredata.UserOrganization{
|
||||
UserID: userID,
|
||||
OrganizationID: organizationID,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
|
||||
return s.pg.WithTx(
|
||||
return s.pg.WithConn(
|
||||
ctx,
|
||||
func(tx pg.Conn) error {
|
||||
_, err := tx.Exec(ctx, q, args)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to add user to organization: %w", err)
|
||||
}
|
||||
return nil
|
||||
return uo.Insert(ctx, tx)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
// GetUserIDFromContext gets the user ID from the context
|
||||
func (s Service) GetUserIDFromContext(ctx context.Context) (gid.GID, error) {
|
||||
// Get the session ID from the context
|
||||
sessionID, ok := ctx.Value("session_id").(gid.GID)
|
||||
if !ok {
|
||||
return gid.GID{}, fmt.Errorf("no session ID in context")
|
||||
}
|
||||
|
||||
// Get the session
|
||||
session, err := s.GetSession(ctx, sessionID)
|
||||
if err != nil {
|
||||
return gid.GID{}, fmt.Errorf("failed to get session: %w", err)
|
||||
}
|
||||
|
||||
return session.UserID, nil
|
||||
}
|
||||
|
||||
// UpdateSession updates a session in the database
|
||||
func (s Service) UpdateSession(
|
||||
ctx context.Context,
|
||||
session *coredata.Session,
|
||||
|
||||
Reference in New Issue
Block a user