package auth import ( "context" "crypto/rand" "encoding/hex" "errors" "strings" "time" "github.com/google/uuid" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) var ( ErrInvalidCredentials = errors.New("invalid credentials") ErrMustSetPassword = errors.New("password_not_set") ErrInviteInvalid = errors.New("invite invalid or expired") ErrPasswordAlreadySet = errors.New("password already set") ErrUserExists = errors.New("user already exists") ) type Service struct { Pool *pgxpool.Pool } type User struct { ID uuid.UUID `json:"id"` Email string `json:"email"` Name *string `json:"name,omitempty"` MustSetPassword bool `json:"must_set_password"` IsPlatformAdmin bool `json:"is_platform_admin"` StaffRole *string `json:"staff_role,omitempty"` IsActive bool `json:"is_active"` } type Membership struct { ID uuid.UUID `json:"id"` CompanyID uuid.UUID `json:"company_id"` UserID uuid.UUID `json:"user_id"` Role string `json:"role"` Status string `json:"status"` } type Company struct { ID uuid.UUID `json:"id"` Name string `json:"name"` } type RegisterInput struct { Email string Password string Name string CompanyName string } type LoginResult struct { User User `json:"user"` CompanyID uuid.UUID `json:"company_id"` Companies []Company `json:"companies"` // ProvisionCompanyIDs are newly created tenants that still need Free-plan billing setup. ProvisionCompanyIDs []uuid.UUID `json:"-"` } func (s *Service) Register(ctx context.Context, in RegisterInput) (LoginResult, error) { email := strings.ToLower(strings.TrimSpace(in.Email)) if email == "" || in.Password == "" || strings.TrimSpace(in.CompanyName) == "" { return LoginResult{}, ErrRegisterFieldsRequired } hash, err := HashPassword(in.Password) if err != nil { return LoginResult{}, err } tx, err := s.Pool.Begin(ctx) if err != nil { return LoginResult{}, err } defer tx.Rollback(ctx) var existing uuid.UUID err = tx.QueryRow(ctx, `SELECT id FROM users WHERE email = $1`, email).Scan(&existing) if err == nil { return LoginResult{}, ErrUserExists } if !errors.Is(err, pgx.ErrNoRows) { return LoginResult{}, err } var userID uuid.UUID var name *string if strings.TrimSpace(in.Name) != "" { n := strings.TrimSpace(in.Name) name = &n } err = tx.QueryRow(ctx, ` INSERT INTO users (email, name, password_hash, must_set_password) VALUES ($1, $2, $3, false) RETURNING id`, email, name, hash).Scan(&userID) if err != nil { return LoginResult{}, err } companyID, created, err := ensureOwnedCompanyTx(ctx, tx, userID, strings.TrimSpace(in.CompanyName)) if err != nil { return LoginResult{}, err } if !created { return LoginResult{}, errors.New("owned company unexpectedly already existed during register") } if err := tx.Commit(ctx); err != nil { return LoginResult{}, err } user, err := s.GetUser(ctx, userID) if err != nil { return LoginResult{}, err } return LoginResult{ User: user, CompanyID: companyID, Companies: []Company{{ID: companyID, Name: strings.TrimSpace(in.CompanyName)}}, ProvisionCompanyIDs: []uuid.UUID{companyID}, }, nil } func (s *Service) Login(ctx context.Context, email, password string) (LoginResult, error) { email = strings.ToLower(strings.TrimSpace(email)) var ( user User hash *string ) err := s.Pool.QueryRow(ctx, ` SELECT id, email, name, password_hash, must_set_password, is_platform_admin, staff_role, is_active FROM users WHERE email = $1`, email).Scan( &user.ID, &user.Email, &user.Name, &hash, &user.MustSetPassword, &user.IsPlatformAdmin, &user.StaffRole, &user.IsActive, ) if isUndefinedColumn(err) { err = s.Pool.QueryRow(ctx, ` SELECT id, email, name, password_hash, must_set_password, is_platform_admin, is_active FROM users WHERE email = $1`, email).Scan( &user.ID, &user.Email, &user.Name, &hash, &user.MustSetPassword, &user.IsPlatformAdmin, &user.IsActive, ) } if errors.Is(err, pgx.ErrNoRows) { return LoginResult{}, ErrInvalidCredentials } if err != nil { return LoginResult{}, err } if !user.IsActive { return LoginResult{}, ErrInvalidCredentials } // Migrated / invite-pending accounts have no usable password until accept-invite or set-password. if user.MustSetPassword || hash == nil || *hash == "" { if user.MustSetPassword { return LoginResult{}, ErrMustSetPassword } return LoginResult{}, ErrInvalidCredentials } ok, err := VerifyPassword(*hash, password) if err != nil || !ok { return LoginResult{}, ErrInvalidCredentials } companies, err := s.ListUserCompanies(ctx, user.ID) if err != nil { return LoginResult{}, err } var companyID uuid.UUID if len(companies) > 0 { companyID = companies[0].ID } _, _ = s.Pool.Exec(ctx, `UPDATE users SET last_login_at = now(), updated_at = now() WHERE id = $1`, user.ID) return LoginResult{User: user, CompanyID: companyID, Companies: companies}, nil } func (s *Service) AcceptInvite(ctx context.Context, token, password, name string) (LoginResult, error) { token = strings.TrimSpace(token) if token == "" { return LoginResult{}, ErrInviteInvalid } var ( inviteID, companyID uuid.UUID email, role string expiresAt time.Time acceptedAt *time.Time ) // Prefer hashed lookup (at-rest); also accept legacy plaintext rows until they expire. tokenHash := HashInviteToken(token) err := s.Pool.QueryRow(ctx, ` SELECT id, company_id, email, role, expires_at, accepted_at FROM invites WHERE token = $1 OR token = $2 ORDER BY CASE WHEN token = $1 THEN 0 ELSE 1 END LIMIT 1`, tokenHash, token).Scan( &inviteID, &companyID, &email, &role, &expiresAt, &acceptedAt, ) if errors.Is(err, pgx.ErrNoRows) || (acceptedAt != nil) || time.Now().After(expiresAt) { return LoginResult{}, ErrInviteInvalid } if err != nil { return LoginResult{}, err } tx, err := s.Pool.Begin(ctx) if err != nil { return LoginResult{}, err } defer tx.Rollback(ctx) var userID uuid.UUID var existingHash *string var mustSet bool err = tx.QueryRow(ctx, ` SELECT id, password_hash, must_set_password FROM users WHERE email = $1`, strings.ToLower(email)).Scan(&userID, &existingHash, &mustSet) displayName := strings.TrimSpace(name) if errors.Is(err, pgx.ErrNoRows) { hash, herr := HashPassword(password) if herr != nil { return LoginResult{}, herr } var n *string if displayName != "" { n = &displayName } err = tx.QueryRow(ctx, ` INSERT INTO users (email, name, password_hash, must_set_password) VALUES ($1, $2, $3, false) RETURNING id`, strings.ToLower(email), n, hash).Scan(&userID) if err != nil { return LoginResult{}, err } } else if err != nil { return LoginResult{}, err } else { var activeMemberships int64 if err := tx.QueryRow(ctx, ` SELECT count(*) FROM memberships WHERE user_id = $1 AND status = 'active'`, userID).Scan(&activeMemberships); err != nil { return LoginResult{}, err } // New password when must_set, or orphan recovery (removed from all companies; no SMTP reset). if mustSet || activeMemberships == 0 { hash, herr := HashPassword(password) if herr != nil { return LoginResult{}, herr } _, err = tx.Exec(ctx, ` UPDATE users SET password_hash = $2, must_set_password = false, updated_at = now() WHERE id = $1`, userID, hash) if err != nil { return LoginResult{}, err } } else { // Active elsewhere: keep password; invitee must prove ownership. hash := "" if existingHash != nil { hash = *existingHash } ok, verr := VerifyPassword(hash, password) if verr != nil || !ok { return LoginResult{}, ErrInvalidCredentials } } } // Always ensure a personal owned workspace; invite also joins the inviting company. ownedName := defaultOwnedCompanyName(email, displayName) ownedID, ownedCreated, err := ensureOwnedCompanyTx(ctx, tx, userID, ownedName) if err != nil { return LoginResult{}, err } _, err = tx.Exec(ctx, ` INSERT INTO memberships (company_id, user_id, role, status) VALUES ($1, $2, $3, 'active') ON CONFLICT (company_id, user_id) DO UPDATE SET role = CASE WHEN memberships.role = 'admin' THEN memberships.role ELSE EXCLUDED.role END, status = 'active', updated_at = now()`, companyID, userID, role) if err != nil { return LoginResult{}, err } ct, err := tx.Exec(ctx, ` UPDATE invites SET accepted_at = now() WHERE id = $1 AND accepted_at IS NULL`, inviteID) if err != nil { return LoginResult{}, err } if ct.RowsAffected() == 0 { return LoginResult{}, ErrInviteInvalid } if err := tx.Commit(ctx); err != nil { return LoginResult{}, err } user, err := s.GetUser(ctx, userID) if err != nil { return LoginResult{}, err } companies, err := s.ListUserCompanies(ctx, userID) if err != nil { return LoginResult{}, err } res := LoginResult{User: user, CompanyID: companyID, Companies: companies} if ownedCreated { res.ProvisionCompanyIDs = []uuid.UUID{ownedID} } return res, nil } func (s *Service) SetPassword(ctx context.Context, userID uuid.UUID, password string) error { hash, err := HashPassword(password) if err != nil { return err } // Only users flagged must_set_password may set via token/session bootstrap. // This also makes HMAC set-password tokens single-use after success. ct, err := s.Pool.Exec(ctx, ` UPDATE users SET password_hash = $2, must_set_password = false, session_version = session_version + 1, updated_at = now() WHERE id = $1 AND must_set_password = true`, userID, hash) if isUndefinedColumn(err) { ct, err = s.Pool.Exec(ctx, ` UPDATE users SET password_hash = $2, must_set_password = false, updated_at = now() WHERE id = $1 AND must_set_password = true`, userID, hash) } if err != nil { return err } if ct.RowsAffected() == 0 { var exists bool _ = s.Pool.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM users WHERE id = $1)`, userID).Scan(&exists) if exists { return ErrPasswordAlreadySet } return ErrUserNotFound } return nil } // ChangePassword verifies the current password then sets a new one (in-app Settings). // Bumps session_version so other sessions are revoked; callers must re-stamp the cookie. func (s *Service) ChangePassword(ctx context.Context, userID uuid.UUID, currentPassword, newPassword string) error { var ( hash string mustSetPassword bool isActive bool ) err := s.Pool.QueryRow(ctx, ` SELECT password_hash, must_set_password, is_active FROM users WHERE id = $1`, userID).Scan(&hash, &mustSetPassword, &isActive) if errors.Is(err, pgx.ErrNoRows) { return ErrUserNotFound } if err != nil { return err } if !isActive { return ErrUserNotFound } if mustSetPassword { return ErrMustSetPassword } ok, err := VerifyPassword(hash, currentPassword) if err != nil { return err } if !ok { return ErrInvalidCredentials } newHash, err := HashPassword(newPassword) if err != nil { return err } ct, err := s.Pool.Exec(ctx, ` UPDATE users SET password_hash = $2, must_set_password = false, session_version = session_version + 1, updated_at = now() WHERE id = $1 AND is_active = true AND must_set_password = false`, userID, newHash) if isUndefinedColumn(err) { ct, err = s.Pool.Exec(ctx, ` UPDATE users SET password_hash = $2, must_set_password = false, updated_at = now() WHERE id = $1 AND is_active = true AND must_set_password = false`, userID, newHash) } if err != nil { return err } if ct.RowsAffected() == 0 { return ErrUserNotFound } return nil } // ForceSetPassword sets a password regardless of must_set_password (local/admin bootstrap). func (s *Service) ForceSetPassword(ctx context.Context, userID uuid.UUID, password string) error { hash, err := HashPassword(password) if err != nil { return err } ct, err := s.Pool.Exec(ctx, ` UPDATE users SET password_hash = $2, must_set_password = false, updated_at = now() WHERE id = $1 AND is_active = true`, userID, hash) if err != nil { return err } if ct.RowsAffected() == 0 { return ErrUserNotFound } return nil } func (s *Service) GetUser(ctx context.Context, id uuid.UUID) (User, error) { var u User err := s.Pool.QueryRow(ctx, ` SELECT id, email, name, must_set_password, is_platform_admin, staff_role, is_active FROM users WHERE id = $1`, id).Scan( &u.ID, &u.Email, &u.Name, &u.MustSetPassword, &u.IsPlatformAdmin, &u.StaffRole, &u.IsActive, ) if isUndefinedColumn(err) { err = s.Pool.QueryRow(ctx, ` SELECT id, email, name, must_set_password, is_platform_admin, is_active FROM users WHERE id = $1`, id).Scan( &u.ID, &u.Email, &u.Name, &u.MustSetPassword, &u.IsPlatformAdmin, &u.IsActive, ) } return u, err } func (s *Service) ListUserCompanies(ctx context.Context, userID uuid.UUID) ([]Company, error) { // Prefer Platform Demo sandbox when present, then richest tenant (products/feeds). // A1 Slovenija wins remaining ties; accept old Local Demo Co rename as A1 alias. const a1LegacyCompanyID = "97e1a309-3d23-4aa2-b518-8e8d7afdfec7" rows, err := s.Pool.Query(ctx, ` SELECT c.id, c.name FROM memberships m JOIN companies c ON c.id = m.company_id WHERE m.user_id = $1 AND m.status = 'active' ORDER BY CASE WHEN lower(c.name) IN ('platform demo', 'demo') THEN 0 ELSE 1 END, (SELECT COUNT(*) FROM processed_products p WHERE p.company_id = c.id) DESC, (SELECT COUNT(*) FROM input_feeds f WHERE f.company_id = c.id) DESC, CASE WHEN lower(c.name) = 'a1 slovenija' THEN 0 WHEN lower(c.name) = 'local demo co' THEN 0 WHEN lower(COALESCE(c.legacy_company_id, '')) = lower($2) THEN 0 ELSE 1 END, c.name`, userID, a1LegacyCompanyID) if err != nil { return nil, err } defer rows.Close() var out []Company for rows.Next() { var c Company if err := rows.Scan(&c.ID, &c.Name); err != nil { return nil, err } out = append(out, c) } return out, rows.Err() } // ListCompanies returns tenants for platform staff company switching (id + name). // Ordered with A1 / demo sandboxes first, then name. Caps at limit (default 500). func (s *Service) ListCompanies(ctx context.Context, limit int) ([]Company, error) { if s == nil || s.Pool == nil { return nil, errors.New("auth service unavailable") } if limit <= 0 || limit > 2000 { limit = 500 } const a1LegacyCompanyID = "97e1a309-3d23-4aa2-b518-8e8d7afdfec7" rows, err := s.Pool.Query(ctx, ` SELECT c.id, c.name FROM companies c ORDER BY CASE WHEN lower(c.name) = 'a1 slovenija' THEN 0 WHEN lower(c.name) = 'local demo co' THEN 0 WHEN lower(COALESCE(c.legacy_company_id, '')) = lower($1) THEN 0 WHEN lower(c.name) IN ('platform demo', 'demo') THEN 1 ELSE 2 END, c.name LIMIT $2`, a1LegacyCompanyID, limit) if err != nil { return nil, err } defer rows.Close() out := make([]Company, 0) for rows.Next() { var c Company if err := rows.Scan(&c.ID, &c.Name); err != nil { return nil, err } out = append(out, c) } return out, rows.Err() } // CompanyByID loads a single company or ErrNotCompanyMember-style miss via pgx.ErrNoRows. func (s *Service) CompanyByID(ctx context.Context, companyID uuid.UUID) (Company, error) { if s == nil || s.Pool == nil { return Company{}, errors.New("auth service unavailable") } var c Company err := s.Pool.QueryRow(ctx, `SELECT id, name FROM companies WHERE id = $1`, companyID).Scan(&c.ID, &c.Name) if errors.Is(err, pgx.ErrNoRows) { return Company{}, ErrCompanyNotFound } return c, err } func (s *Service) EnsureMembership(ctx context.Context, userID, companyID uuid.UUID) (Membership, error) { if s == nil || s.Pool == nil { return Membership{}, ErrNotCompanyMember } var m Membership err := s.Pool.QueryRow(ctx, ` SELECT id, company_id, user_id, role, status FROM memberships WHERE user_id = $1 AND company_id = $2 AND status = 'active'`, userID, companyID).Scan(&m.ID, &m.CompanyID, &m.UserID, &m.Role, &m.Status) if errors.Is(err, pgx.ErrNoRows) { return Membership{}, ErrNotCompanyMember } return m, err } func RandomToken(n int) (string, error) { b := make([]byte, n) if _, err := rand.Read(b); err != nil { return "", err } return hex.EncodeToString(b), nil }