#!/usr/bin/python2
# vim: tabstop=4 shiftwidth=4 softtabstop=4

# Copyright 2016-2017 Lenovo
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from collections import deque
import optparse
import os
import select
import signal
import subprocess
import sys

try:
    signal.signal(signal.SIGPIPE, signal.SIG_DFL)
except AttributeError:
    pass
path = os.path.dirname(os.path.realpath(__file__))
path = os.path.realpath(os.path.join(path, '..', 'lib', 'python'))
if path.startswith('/opt'):
    sys.path.append(path)

import confluent.client as client
import confluent.sortutil as sortutil


def run():

    argparser = optparse.OptionParser(
        usage="Usage: %prog [options] noderange commandexpression",
        epilog="Expressions are the same as in attributes, e.g. "
               "'ipmitool -H {hardwaremanagement.manager}' will be expanded.")
    argparser.add_option('-f', '-c', '--count', type='int', default=168,
                         help='Number of commands to run at a time')
    argparser.add_option('-n', '--nonodeprefix', action='store_true',
                         help='Do not prefix output with node names')
    argparser.add_option('-p', '--port', type='int', default=0,
                         help='Specify a custom port for ssh')
    argparser.add_option('-m', '--maxnodes', type='int',
                     help='Specify a maximum number of '
                          'nodes to run remote ssh command to, '
                          'prompting if over the threshold')
    # among other things, FD_SETSIZE limits.  Besides, spawning too many
    # processes can be unkind for the unaware on memory pressure and such...
    argparser.disable_interspersed_args()
    (options, args) = argparser.parse_args()
    if len(args) < 2:
        argparser.print_help()
        sys.exit(1)
    client.check_globbing(args[0])
    concurrentprocs = options.count
    c = client.Command()
    cmdstr = " ".join(args[1:])

    currprocs = 0
    all = set([])
    pipedesc = {}
    pendingexecs = deque()
    exitcode = 0

    c.stop_if_noderange_over(args[0], options.maxnodes)
    for exp in c.create('/noderange/{0}/attributes/expression'.format(args[0]),
                        {'expression': cmdstr}):
        if 'error' in exp:
            sys.stderr.write(exp['error'] + '\n')
            exitcode |= exp.get('errorcode', 1)
        ex = exp.get('databynode', ())
        for node in ex:
            cmd = ex[node]['value']
            if not isinstance(cmd, str) and not isinstance(cmd, bytes):
                cmd = cmd.encode('utf-8')
            if options.port:
                cmdv = ['ssh', '-p', '{0}'.format(options.port), node, cmd]
            else:
                cmdv = ['ssh', node, cmd]
            if currprocs < concurrentprocs:
                currprocs += 1
                run_cmdv(node, cmdv, all, pipedesc)
            else:
                pendingexecs.append((node, cmdv))
    if not all or exitcode:
        sys.exit(exitcode)
    rdy, _, _ = select.select(all, [], [], 10)
    while all:
        pernodeout = {}
        for r in rdy:
            desc = pipedesc[r]
            node = desc['node']
            data = True
            while data and select.select([r], [], [], 0)[0]:
                data = r.readline()
                if data:
                    if desc['type'] == 'stdout':
                        if node not in pernodeout:
                            pernodeout[node] = []
                        pernodeout[node].append(data)
                    else:
                        data = client.stringify(data)
                        if options.nonodeprefix:
                            sys.stderr.write(data)
                        else:
                            sys.stderr.write('{0}: {1}'.format(node, data))
                        sys.stderr.flush()
                else:
                    pop = desc['popen']
                    ret = pop.poll()
                    if ret is not None:
                        exitcode = exitcode | ret
                        all.discard(r)
                        r.close()
                        if desc['type'] == 'stdout' and pendingexecs:
                            node, cmdv = pendingexecs.popleft()
                            run_cmdv(node, cmdv, all, pipedesc)
        for node in sortutil.natural_sort(pernodeout):
            for line in pernodeout[node]:
                line = client.stringify(line)
                if options.nonodeprefix:
                    sys.stdout.write(line)
                else:
                    sys.stdout.write('{0}: {1}'.format(node, line))
        sys.stdout.flush()
        if all:
            rdy, _, _ = select.select(all, [], [], 10)
    sys.exit(exitcode)


def run_cmdv(node, cmdv, all, pipedesc):
    nopen = subprocess.Popen(
        cmdv, stdout=subprocess.PIPE, stderr=subprocess.PIPE)
    pipedesc[nopen.stdout] = {'node': node, 'popen': nopen,
                              'type': 'stdout'}
    pipedesc[nopen.stderr] = {'node': node, 'popen': nopen,
                              'type': 'stderr'}
    all.add(nopen.stdout)
    all.add(nopen.stderr)


if __name__ == '__main__':
    run()