#!/usr/bin/env 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.

# Runs from cron as somon, a docker group member, and emits influx line protocol for telegraf to
# read. This exists so so-telegraf does not need the docker socket; the container reads the output
# file instead.
# Measurement/tag/field names match telegraf's inputs.docker plugin because the InfluxDB
# dashboards query them directly. Which fields are emitted is set per stat in SOC; see
# telegraf.container_stats in defaults.yaml.

import json

SETTINGS = json.loads('''{{ CONTAINER_STATS | tojson }}''')
{% raw %}
import os
import re
import subprocess
import sys
from datetime import datetime, timezone

CGROUP_ROOT = '/sys/fs/cgroup'
NANOSEC = 10**9
# cgroup and docker report cpu time in microseconds; inputs.docker published nanoseconds
USEC_TO_NSEC = 1000


def want(group, field):
  return bool(SETTINGS.get(group, {}).get(field))


def wants_any(group, fields):
  return any(want(group, field) for field in fields)


def docker(args):
  proc = subprocess.run(['docker'] + args, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, encoding='utf-8')
  if proc.returncode != 0:
    print('Container system error; unable to query docker', file=sys.stderr)
    sys.exit(1)
  return proc.stdout


def escape_tag(value):
  return value.replace(',', '\\,').replace(' ', '\\ ').replace('=', '\\=')


def quote(value):
  return '"{}"'.format(str(value).replace('\\', '\\\\').replace('"', '\\"'))


def integer(value):
  return '{}i'.format(int(value))


def unsigned(value):
  # inputs.docker published the cgroup and network counters as uint64; influx stores u and i
  # as different field types, so matching it keeps historical series readable
  return '{}u'.format(max(int(value), 0))


def read_text(path):
  try:
    with open(path) as handle:
      return handle.read()
  except OSError:
    return ''


def read_pairs(path):
  # cgroup files such as memory.stat and cpu.stat are "key value" per line
  values = {}
  for line in read_text(path).splitlines():
    parts = line.split()
    if len(parts) == 2:
      try:
        values[parts[0]] = int(parts[1])
      except ValueError:
        pass
  return values


def read_value(path):
  raw = read_text(path).strip()
  try:
    return int(raw)
  except ValueError:
    return None


def cgroup_path(pid):
  # "0::/system.slice/docker-<id>.scope" on cgroup v2
  for line in read_text('/proc/{}/cgroup'.format(pid)).splitlines():
    parts = line.split(':', 2)
    if len(parts) == 3 and parts[1] == '':
      return CGROUP_ROOT + parts[2]
  return None


def host_memory_total():
  match = re.search(r'MemTotal:\s+(\d+) kB', read_text('/proc/meminfo'))
  return int(match.group(1)) * 1024 if match else 0


def host_cpu_nanoseconds():
  # inputs.docker's usage_system is the host-wide cpu time the daemon reads from /proc/stat
  for line in read_text('/proc/stat').splitlines():
    if line.startswith('cpu '):
      ticks = sum(int(value) for value in line.split()[1:])
      return int(ticks * NANOSEC / os.sysconf('SC_CLK_TCK'))
  return 0


def parse_image(image):
  # mirrors telegraf's internal/docker ParseImage so the tags match what inputs.docker emitted
  domain = ''
  remainder = image
  if '/' in image:
    head, _, tail = image.partition('/')
    if '.' in head or ':' in head or head == 'localhost':
      domain, remainder = head + '/', tail
  if ':' in remainder:
    name, _, version = remainder.rpartition(':')
    return domain + name, version
  return domain + remainder, 'unknown'


def to_nanoseconds(stamp):
  # docker emits 9 fractional digits; fromisoformat takes at most 6 before python 3.11
  stamp = stamp.rstrip('Z')
  if '.' in stamp:
    whole, _, frac = stamp.partition('.')
    stamp = whole + '.' + frac[:6]
  try:
    parsed = datetime.fromisoformat(stamp).replace(tzinfo=timezone.utc)
  except ValueError:
    return None
  if parsed.year <= 1:
    return None
  return int(parsed.timestamp() * NANOSEC)


def percent(value):
  try:
    return float(value.strip().rstrip('%'))
  except ValueError:
    return 0.0


def net_counters(pid):
  # summed across interfaces except loopback, matching the daemon's per-container totals
  rows = read_text('/proc/{}/net/dev'.format(pid)).splitlines()[2:]
  if not rows:
    return None
  names = ['rx_bytes', 'rx_packets', 'rx_errors', 'rx_dropped']
  totals = dict.fromkeys(names + ['tx_bytes', 'tx_packets', 'tx_errors', 'tx_dropped'], 0)
  found = False
  for row in rows:
    iface, _, rest = row.partition(':')
    if iface.strip() == 'lo':
      continue
    columns = rest.split()
    if len(columns) < 12:
      continue
    found = True
    for index, name in enumerate(names):
      totals[name] += int(columns[index])
    for index, name in enumerate(['tx_bytes', 'tx_packets', 'tx_errors', 'tx_dropped']):
      totals[name] += int(columns[8 + index])
  return totals if found else None


def blkio_counters(path):
  # inputs.docker emitted the device=total row unconditionally, so a container that has done
  # no block io reports zeros rather than dropping out of the measurement entirely
  totals = {'io_service_bytes_recursive_read': 0, 'io_service_bytes_recursive_write': 0}
  rows = read_text(path + '/io.stat').splitlines()
  for row in rows:
    for token in row.split()[1:]:
      key, _, value = token.partition('=')
      try:
        if key == 'rbytes':
          totals['io_service_bytes_recursive_read'] += int(value)
        elif key == 'wbytes':
          totals['io_service_bytes_recursive_write'] += int(value)
      except ValueError:
        pass
  return totals


def collect_stats():
  stats = {}
  for line in docker(['stats', '--no-stream', '--format', '{{json .}}']).splitlines():
    line = line.strip()
    if not line:
      continue
    try:
      entry = json.loads(line)
    except json.JSONDecodeError:
      continue
    stats[entry.get('Name', '')] = entry
  return stats


def collect_inspect():
  ids = docker(['ps', '-aq']).split()
  if not ids:
    return []
  try:
    return json.loads(docker(['inspect'] + ids))
  except json.JSONDecodeError:
    return []


def collect_info():
  try:
    return json.loads(docker(['info', '--format', '{{json .}}']))
  except json.JSONDecodeError:
    return {}


class Emitter:
  def __init__(self):
    self.lines = []

  def add(self, measurement, tags, fields):
    if not fields:
      return
    tagset = ','.join('{}={}'.format(key, escape_tag(value)) for key, value in tags)
    body = ','.join('{}={}'.format(key, fields[key]) for key in sorted(fields))
    self.lines.append('{},{} {}'.format(measurement, tagset, body))


def main():
  cpu_extra = ['usage_total', 'usage_in_usermode', 'usage_in_kernelmode',
               'throttling_periods', 'throttling_throttled_periods', 'throttling_throttled_time']
  mem_extra = ['usage', 'limit', 'max_usage', 'active_anon', 'active_file', 'inactive_anon',
               'inactive_file', 'unevictable', 'pgfault', 'pgmajfault']
  net_fields = ['rx_bytes', 'rx_packets', 'rx_errors', 'rx_dropped',
                'tx_bytes', 'tx_packets', 'tx_errors', 'tx_dropped']
  engine_fields = ['n_containers', 'n_containers_running', 'n_containers_stopped',
                   'n_containers_paused', 'n_images', 'n_cpus', 'n_goroutines',
                   'n_used_file_descriptors', 'n_listener_events']

  identity = want('tags', 'identity')
  need_stats = want('cpu', 'usage_percent')
  need_cgroup_cpu = wants_any('cpu', cpu_extra)
  need_cgroup_mem = wants_any('mem', mem_extra + ['usage_percent'])
  need_blkio = wants_any('blkio', ['io_service_bytes_recursive_read',
                                   'io_service_bytes_recursive_write', 'container_id'])
  need_net = wants_any('net', net_fields + ['container_id'])
  need_info = identity or wants_any('engine', engine_fields + ['memory_total'])

  emitter = Emitter()
  stats = collect_stats() if need_stats else {}
  info = collect_info() if need_info else {}
  engine_tags = [('engine_host', info.get('Name', '')),
                 ('server_version', info.get('ServerVersion', ''))]
  system_ns = host_cpu_nanoseconds() if want('cpu', 'usage_system') else 0
  mem_total = host_memory_total() if wants_any('mem', ['limit', 'usage_percent']) else 0

  if info:
    counts = {'n_containers': 'Containers', 'n_containers_running': 'ContainersRunning',
              'n_containers_stopped': 'ContainersStopped', 'n_containers_paused': 'ContainersPaused',
              'n_images': 'Images', 'n_cpus': 'NCPU', 'n_goroutines': 'NGoroutines',
              'n_used_file_descriptors': 'NFd', 'n_listener_events': 'NEventsListener'}
    fields = {name: integer(info[key]) for name, key in counts.items()
              if want('engine', name) and key in info}
    emitter.add('docker', engine_tags, fields)
    # inputs.docker published memory_total as its own point
    if want('engine', 'memory_total') and 'MemTotal' in info:
      emitter.add('docker', engine_tags, {'memory_total': integer(info['MemTotal'])})

  for container in collect_inspect():
    name = container.get('Name', '').lstrip('/')
    if not name:
      continue
    state = container.get('State', {})
    status = state.get('Status', 'unknown')
    pid = state.get('Pid')
    container_id = container.get('Id', '')
    tags = [('container_name', name), ('container_status', status)]
    if identity:
      image, version = parse_image(container.get('Config', {}).get('Image', ''))
      tags += [('container_image', image), ('container_version', version)] + engine_tags

    status_fields = {}
    if want('status', 'uptime_ns') or want('status', 'started_at') or want('status', 'finished_at'):
      started = to_nanoseconds(state.get('StartedAt', ''))
      finished = to_nanoseconds(state.get('FinishedAt', ''))
      if started is not None:
        if want('status', 'started_at'):
          status_fields['started_at'] = integer(started)
        if want('status', 'uptime_ns'):
          end = finished if finished is not None and finished >= started else int(
            datetime.now(timezone.utc).timestamp() * NANOSEC)
          status_fields['uptime_ns'] = integer(end - started)
      if finished is not None and want('status', 'finished_at'):
        status_fields['finished_at'] = integer(finished)
    if want('status', 'oomkilled'):
      status_fields['oomkilled'] = 'true' if state.get('OOMKilled') else 'false'
    if want('status', 'pid'):
      status_fields['pid'] = integer(pid or 0)
    if want('status', 'exitcode'):
      status_fields['exitcode'] = integer(state.get('ExitCode', 0))
    if want('status', 'restart_count'):
      status_fields['restart_count'] = integer(container.get('RestartCount', 0))
    if want('status', 'container_id'):
      status_fields['container_id'] = quote(container_id)
    emitter.add('docker_container_status', tags, status_fields)

    health = state.get('Health')
    if health:
      health_fields = {}
      if want('health', 'health_status'):
        health_fields['health_status'] = quote(health.get('Status', ''))
      if want('health', 'failing_streak'):
        health_fields['failing_streak'] = integer(health.get('FailingStreak', 0))
      emitter.add('docker_container_health', tags, health_fields)

    if status != 'running' or not pid:
      continue
    path = cgroup_path(pid)
    entry = stats.get(name, {})

    cpu_fields = {}
    if want('cpu', 'usage_percent'):
      cpu_fields['usage_percent'] = percent(entry.get('CPUPerc', '0%'))
    if want('cpu', 'usage_system'):
      cpu_fields['usage_system'] = unsigned(system_ns)
    if want('cpu', 'container_id'):
      cpu_fields['container_id'] = quote(container_id)
    if need_cgroup_cpu and path:
      cpu = read_pairs(path + '/cpu.stat')
      mapping = {'usage_total': 'usage_usec', 'usage_in_usermode': 'user_usec',
                 'usage_in_kernelmode': 'system_usec',
                 'throttling_periods': 'nr_periods',
                 'throttling_throttled_periods': 'nr_throttled',
                 'throttling_throttled_time': 'throttled_usec'}
      for field, key in mapping.items():
        if want('cpu', field) and key in cpu:
          scale = 1 if field.startswith('throttling_') and field != 'throttling_throttled_time' else USEC_TO_NSEC
          cpu_fields[field] = unsigned(cpu[key] * scale)
    emitter.add('docker_container_cpu', tags + [('cpu', 'cpu-total')], cpu_fields)

    mem_fields = {}
    if want('mem', 'container_id'):
      mem_fields['container_id'] = quote(container_id)
    if need_cgroup_mem and path:
      memory = read_pairs(path + '/memory.stat')
      for field in ['active_anon', 'active_file', 'inactive_anon', 'inactive_file',
                    'unevictable', 'pgfault', 'pgmajfault']:
        if want('mem', field) and field in memory:
          mem_fields[field] = unsigned(memory[field])
      current = read_value(path + '/memory.current')
      # inputs.docker reports usage net of reclaimable page cache
      usage = max(current - memory.get('inactive_file', 0), 0) if current is not None else None
      raw_limit = read_value(path + '/memory.max')
      limit = raw_limit if raw_limit is not None else mem_total
      if want('mem', 'usage') and usage is not None:
        mem_fields['usage'] = unsigned(usage)
      if want('mem', 'limit'):
        mem_fields['limit'] = unsigned(limit)
      if want('mem', 'usage_percent') and usage is not None:
        # same ratio inputs.docker computes, from the same two values
        mem_fields['usage_percent'] = usage / limit * 100.0 if limit else 0.0
      if want('mem', 'max_usage'):
        peak = read_value(path + '/memory.peak')
        if peak is not None:
          mem_fields['max_usage'] = unsigned(peak)
    emitter.add('docker_container_mem', tags, mem_fields)

    # inputs.docker emitted nothing for host-network containers; its Networks map was empty
    if need_net and container.get('HostConfig', {}).get('NetworkMode', '') != 'host':
      counters = net_counters(pid)
      if counters is not None:
        net_out = {field: unsigned(counters[field]) for field in net_fields if want('net', field)}
        if want('net', 'container_id'):
          net_out['container_id'] = quote(container_id)
        emitter.add('docker_container_net', tags + [('network', 'total')], net_out)

    if need_blkio and path:
      counters = blkio_counters(path)
      if counters is not None:
        blkio_out = {field: unsigned(value) for field, value in counters.items() if want('blkio', field)}
        if want('blkio', 'container_id'):
          blkio_out['container_id'] = quote(container_id)
        emitter.add('docker_container_blkio', tags + [('device', 'total')], blkio_out)

  print('\n'.join(emitter.lines))


if __name__ == '__main__':
  main()
{% endraw %}
