#!/opt/saltstack/salt/bin/python3

# Copyright Security Onion Solutions LLC and/or licensed to Security Onion Solutions LLC under one
# or more contributor license agreements. Licensed under the Elastic License 2.0 as shown at
# https://securityonion.net/license; you may not use this file except in compliance with the
# Elastic License 2.0.

"""
so-push-drainer
===============

Scheduled drainer for the active-push feature. Runs on the manager every
drain_interval seconds (default 15) via a salt schedule in salt/salt/push_drain_schedule.sls.

For each intent file under /opt/so/state/push_pending/*.json whose last_touch
is older than debounce_seconds, this script:
  * concatenates the actions lists from every ready intent
  * dedupes by (state or __highstate__, tgt, tgt_type)
  * dispatches a single `salt-run state.orchestrate orch.push_batch --async`
    with the deduped actions list passed as pillar kwargs
  * deletes the contributed intent files on successful dispatch
  * records the orchestration jid under /opt/so/state/push_dispatched and, on
    later passes, looks up its result and logs success or per-minion failures

Reactor sls files (push_files, push_pillar) write intents
but never dispatch directly
"""

import fcntl
import glob
import json
import logging
import logging.handlers
import os
import re
import subprocess
import sys
import time

import salt.client

PENDING_DIR = '/opt/so/state/push_pending'
LOCK_FILE = os.path.join(PENDING_DIR, '.lock')
LOG_FILE = '/opt/so/log/salt/so-push-drainer.log'

DISPATCHED_DIR = '/opt/so/state/push_dispatched'

HIGHSTATE_SENTINEL = '__highstate__'

RESULT_CHECK_DELAY = 30
RESULT_MAX_AGE = 7200
RESULT_CHECKS_PER_PASS = 5
TEXT_LIMIT = 500

# salt-run --async reports the jid only in a log line (stderr by default).
JID_RE = re.compile(r'salt/run/(\d{20})')


def _make_logger():
    logger = logging.getLogger('so-push-drainer')
    logger.setLevel(logging.INFO)
    if not logger.handlers:
        os.makedirs(os.path.dirname(LOG_FILE), exist_ok=True)
        handler = logging.handlers.RotatingFileHandler(
            LOG_FILE, maxBytes=5 * 1024 * 1024, backupCount=3,
        )
        handler.setFormatter(logging.Formatter(
            '%(asctime)s | %(levelname)s | %(message)s',
        ))
        logger.addHandler(handler)
    return logger


def _load_push_cfg():
    """Read the salt:auto_apply pillar subtree via salt-call. Returns a dict."""
    caller = salt.client.Caller()
    cfg = caller.cmd('pillar.get', 'salt:auto_apply', {})
    return cfg if isinstance(cfg, dict) else {}


def _read_intent(path, log):
    try:
        with open(path, 'r') as f:
            return json.load(f)
    except (IOError, ValueError) as exc:
        log.warning('cannot read intent %s: %s', path, exc)
        return None
    except Exception:
        log.exception('unexpected error reading %s', path)
        return None


def _dedupe_actions(actions):
    seen = set()
    deduped = []
    for action in actions:
        if not isinstance(action, dict):
            continue
        state_key = HIGHSTATE_SENTINEL if action.get('highstate') else action.get('state')
        tgt = action.get('tgt')
        tgt_type = action.get('tgt_type', 'compound')
        if not state_key or not tgt:
            continue
        key = (state_key, tgt, tgt_type)
        if key in seen:
            continue
        seen.add(key)
        deduped.append(action)
    return deduped


def _dispatch(actions, log):
    pillar_arg = json.dumps({'actions': actions})
    cmd = [
        'salt-run',
        'state.orchestrate',
        'orch.push_batch',
        'pillar={}'.format(pillar_arg),
        '--async',
    ]
    log.info('dispatching: %s', ' '.join(cmd[:3]) + ' pillar=<{} actions>'.format(len(actions)))
    try:
        result = subprocess.run(
            cmd, check=True, capture_output=True, text=True, timeout=60,
        )
    except subprocess.CalledProcessError as exc:
        log.error('dispatch failed (rc=%s): stdout=%s stderr=%s',
                  exc.returncode, exc.stdout, exc.stderr)
        return None
    except subprocess.TimeoutExpired:
        log.error('dispatch timed out after 60s')
        return None
    except Exception:
        log.exception('dispatch raised')
        return None
    match = JID_RE.search('{}\n{}'.format(result.stderr or '', result.stdout or ''))
    if not match:
        log.warning('dispatch accepted but no jid found, result will not be tracked: stderr=%s',
                    _trim(result.stderr))
        return ''
    log.info('dispatch accepted: jid=%s', match.group(1))
    return match.group(1)


def _trim(value):
    text = value if isinstance(value, str) else json.dumps(value, default=str)
    lines = [line.strip() for line in text.splitlines() if line.strip()]
    if 'Traceback (most recent call last):' in text:
        # Keep the lead-in and the raised exception; the frames are noise in a log line.
        lines = [text.split('Traceback (most recent call last):', 1)[0].strip(), lines[-1]]
    text = ' '.join(line for line in lines if line)
    return text if len(text) <= TEXT_LIMIT else text[:TEXT_LIMIT] + '...'


def _unlink(path, log):
    try:
        os.unlink(path)
    except OSError:
        log.exception('failed to remove %s', path)


def _record_dispatch(jid, actions, paths, log):
    record = {'jid': jid, 'dispatched_at': time.time(), 'actions': actions, 'paths': paths}
    path = os.path.join(DISPATCHED_DIR, '{}.json'.format(jid))
    try:
        os.makedirs(DISPATCHED_DIR, exist_ok=True)
        tmp_path = path + '.tmp'
        with open(tmp_path, 'w') as f:
            json.dump(record, f)
        os.rename(tmp_path, path)
    except OSError:
        log.exception('failed to record dispatch %s', jid)


def _lookup_jid(jid, log):
    """Returns the job cache entry for jid, {} while it is still running, or None on error."""
    cmd = ['salt-run', 'jobs.lookup_jid', jid, '--out=json']
    try:
        result = subprocess.run(cmd, check=True, capture_output=True, text=True, timeout=60)
        return json.loads(result.stdout or '{}')
    except (subprocess.CalledProcessError, subprocess.TimeoutExpired, ValueError) as exc:
        log.warning('lookup of jid %s failed: %s', jid, exc)
        return None


def _minion_failure(minion_ret):
    if isinstance(minion_ret, dict):
        return '; '.join(
            '{}: {}'.format(state.get('__id__', state_key), _trim(state.get('comment', '')))
            for state_key, state in minion_ret.items()
            if isinstance(state, dict) and state.get('result') is False
        )
    # A state run rejected before it starts (e.g. another state run is in
    # progress) returns a list of error strings instead of state results.
    if isinstance(minion_ret, (list, str)):
        return _trim(minion_ret)
    return ''


def _step_failures(step):
    if not isinstance(step, dict) or step.get('result') is not False:
        return []
    failures = ['{}: {}'.format(step.get('__id__', step.get('name')), _trim(step.get('comment', '')))]
    changes = step.get('changes')
    minion_rets = changes.get('ret') if isinstance(changes, dict) else None
    if isinstance(minion_rets, dict):
        for minion, minion_ret in minion_rets.items():
            text = _minion_failure(minion_ret)
            if text:
                failures.append('{}: {}'.format(minion, text))
    return failures


def _orch_failures(ret):
    if not isinstance(ret, dict):
        return [_trim(ret)]
    failures = []
    for job in ret.values():
        if not isinstance(job, dict):
            continue
        job_ret = job.get('return')
        data = job_ret.get('data') if isinstance(job_ret, dict) else {}
        if not isinstance(data, dict):
            if data:
                failures.append(_trim(data))
            data = {}
        for steps in data.values():
            if not isinstance(steps, dict):
                failures.append(_trim(steps))
                continue
            for step in steps.values():
                failures.extend(_step_failures(step))
        if job.get('success') is False and not failures:
            failures.append('orchestration reported failure: {}'.format(_trim(job.get('return'))))
    return failures


def _check_dispatched(log, now):
    checked = 0
    for path in sorted(glob.glob(os.path.join(DISPATCHED_DIR, '*.json'))):
        if checked >= RESULT_CHECKS_PER_PASS:
            break
        record = _read_intent(path, log)
        if not isinstance(record, dict) or not record.get('jid'):
            _unlink(path, log)
            continue
        age = now - record.get('dispatched_at', 0)
        if age < RESULT_CHECK_DELAY:
            continue
        checked += 1
        jid = record['jid']
        try:
            if _report_result(record, age, log):
                _unlink(path, log)
        except Exception:
            # Drop the record so one unreadable result can't fail every pass ahead of the drain.
            log.exception('cannot evaluate result for jid=%s; no longer tracking', jid)
            _unlink(path, log)


def _report_result(record, age, log):
    """Logs the outcome of a dispatched push. Returns True once the record is finished with."""
    jid = record['jid']
    paths = record.get('paths', [])
    ret = _lookup_jid(jid, log)
    if not ret:
        if age > RESULT_MAX_AGE:
            log.warning('no result for jid=%s after %ds, no longer tracking; paths=%s', jid, age, paths)
            return True
        return False
    failures = _orch_failures(ret)
    if failures:
        log.error('push failed jid=%s paths=%s; change will be applied at the next scheduled highstate: %s',
                  jid, paths, ' | '.join(failures))
    else:
        log.info('push succeeded jid=%s paths=%s', jid, paths)
    return True


def main():
    log = _make_logger()

    if not os.path.isdir(PENDING_DIR):
        # Nothing to do; reactors create the dir on first use.
        return 0

    try:
        push = _load_push_cfg()
    except Exception:
        log.exception('failed to read salt:auto_apply pillar; aborting drain pass')
        return 1

    if not push.get('enabled', True):
        log.debug('push disabled; exiting')
        return 0

    debounce_seconds = int(push.get('debounce_seconds', 30))

    # Outside the lock: lookups are slow and the reactors take the same lock.
    _check_dispatched(log, time.time())

    os.makedirs(PENDING_DIR, exist_ok=True)
    lock_fd = os.open(LOCK_FILE, os.O_CREAT | os.O_RDWR, 0o644)
    try:
        fcntl.flock(lock_fd, fcntl.LOCK_EX)

        intent_files = [
            p for p in sorted(glob.glob(os.path.join(PENDING_DIR, '*.json')))
            if os.path.basename(p) != '.lock'
        ]
        if not intent_files:
            return 0

        now = time.time()
        ready = []
        skipped = 0
        broken = []
        for path in intent_files:
            intent = _read_intent(path, log)
            if not isinstance(intent, dict):
                broken.append(path)
                continue
            last_touch = intent.get('last_touch', 0)
            if now - last_touch < debounce_seconds:
                skipped += 1
                continue
            ready.append((path, intent))

        for path in broken:
            try:
                os.unlink(path)
            except OSError:
                pass

        if not ready:
            if skipped:
                log.debug('no ready intents (%d still in debounce window)', skipped)
            return 0

        combined_actions = []
        oldest_first_touch = now
        all_paths = []
        for path, intent in ready:
            combined_actions.extend(intent.get('actions', []) or [])
            first = intent.get('first_touch', now)
            if first < oldest_first_touch:
                oldest_first_touch = first
            all_paths.extend(intent.get('paths', []) or [])

        deduped = _dedupe_actions(combined_actions)
        if not deduped:
            log.warning('%d intent(s) had no usable actions; clearing', len(ready))
            for path, _ in ready:
                try:
                    os.unlink(path)
                except OSError:
                    pass
            return 0

        debounce_duration = now - oldest_first_touch
        log.info(
            'draining %d intent(s): %d action(s) after dedupe (raw=%d), '
            'debounce_duration=%.1fs, paths=%s',
            len(ready), len(deduped), len(combined_actions),
            debounce_duration, all_paths[:20],
        )
        for action in deduped:
            log.info('action: %s tgt=%s', 'highstate' if action.get('highstate') else action.get('state'),
                     action.get('tgt'))

        jid = _dispatch(deduped, log)
        if jid is None:
            log.warning('dispatch failed; leaving intent files in place for retry')
            return 1
        if jid:
            _record_dispatch(jid, deduped, all_paths[:20], log)

        for path, _ in ready:
            try:
                os.unlink(path)
            except OSError:
                log.exception('failed to remove drained intent %s', path)

        return 0
    finally:
        try:
            fcntl.flock(lock_fd, fcntl.LOCK_UN)
        finally:
            os.close(lock_fd)


if __name__ == '__main__':
    sys.exit(main())
