package store import ( "context" "database/sql" "atlas9.dev/c/core" "atlas9.dev/c/core/dbi" "atlas9.dev/c/demo/lib/access" "atlas9.dev/c/demo/lib/mfa" ) type SqliteMfaPolicyStore struct { db dbi.DBI guard access.Guard } var _ mfa.PolicyStore = (*SqliteMfaPolicyStore)(nil) func NewSqliteMfaPolicyStore(db dbi.DBI, guard access.Guard) *SqliteMfaPolicyStore { return &SqliteMfaPolicyStore{db: db, guard: guard} } func (s *SqliteMfaPolicyStore) Get(ctx context.Context, tenant core.ID) (bool, error) { if err := s.guard.Check(ctx, mfa.Cap_Mfa_PolicyRead, tenant, ""); err != nil { return false, err } // No row means the tenant has never set a policy: MFA is not required. var required bool err := s.db.QueryRow(ctx, `SELECT required FROM tenant_mfa_policy WHERE tenant = $1`, tenant).Scan(&required) if err == sql.ErrNoRows { return false, nil } return required, err } func (s *SqliteMfaPolicyStore) Set(ctx context.Context, tenant core.ID, required bool) error { if err := s.guard.Check(ctx, mfa.Cap_Mfa_PolicyWrite, tenant, ""); err != nil { return err } _, err := s.db.Exec(ctx, ` INSERT INTO tenant_mfa_policy (tenant, required) VALUES ($1, $2) ON CONFLICT (tenant) DO UPDATE SET required = $2 `, tenant, required) return err } func (s *SqliteMfaPolicyStore) RequiredForUser(ctx context.Context, userID core.ID) (bool, error) { // Answered at login on the system plane, before a session exists. A user is // required to use MFA if any tenant they belong to requires it. if err := s.guard.System(ctx, mfa.Cap_Mfa_Read); err != nil { return false, err } var required bool err := s.db.QueryRow(ctx, ` SELECT EXISTS( SELECT 1 FROM tenant_members m JOIN tenant_mfa_policy p ON p.tenant = m.tenant WHERE m.user = $1 AND p.required = true ) `, userID).Scan(&required) return required, err }