"""Bounded private Docker/host inventory for allocation and scoped diagnostics (Python 3.8)."""
import datetime
import os
import re
import sys
import time

from network_pools import (InventoryFailure, address_exclusions, decode, effective_pools, prefix,
                           resolver_exclusions, route_exclusions, validate_rules)

LIMIT = 1024 * 1024
INFO_FORMAT = ('{"version":{{json .ServerVersion}},"os":{{json .OSType}},'
               '"security":{{json .SecurityOptions}},"pools":{{json .DefaultAddressPools}}}')
HOST_PROBE = '''import json, os, socket, stat, struct, sys
path = sys.argv[1]
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as connection:
    connection.settimeout(3)
    connection.connect(path)
    peer, uid, gid = struct.unpack('3i', connection.getsockopt(socket.SOL_SOCKET, socket.SO_PEERCRED, 12))
record = os.lstat('/run/docker.pid')
with open('/run/docker.pid', encoding='ascii') as source:
    value = source.read(32).strip()
if not value.isdigit() or not 1 < int(value) < 4194304:
    raise SystemExit(1)
pid = int(value)
with open('/proc/%d/comm' % pid, encoding='ascii') as source:
    command = source.read(32).strip()
own = os.stat('/proc/self/ns/net').st_ino
daemon = os.stat('/proc/%d/ns/net' % pid).st_ino
init = os.stat('/proc/1/ns/net').st_ino
rootful = uid == 0 and os.stat('/proc/%d' % pid).st_uid == 0 and peer in (1, pid)
pidfile = stat.S_ISREG(record.st_mode) and record.st_uid == 0 and not record.st_mode & 0o022
print(json.dumps({'local_rootful': rootful and pidfile and command == 'dockerd',
                  'same_namespace': own == daemon == init}))
'''


class Reader:
    def __init__(self, run):
        self.run, self.deadline, self.total = run, time.monotonic() + 45, 0

    def __call__(self, argv, *, parsed=True):
        remaining = self.deadline - time.monotonic()
        if remaining <= 0:
            raise InventoryFailure('inventory_timeout')
        try:
            raw = self.run(argv, seconds=min(10, remaining))
        except Exception as error:
            code = {'command_timeout': 'inventory_timeout', 'output_limit': 'inventory_output_limit'}
            raise InventoryFailure(code.get(getattr(error, 'code', ''), 'inventory_command_failed')) from None
        if not isinstance(raw, bytes) or len(raw) > LIMIT:
            raise InventoryFailure('inventory_output_limit')
        self.total += len(raw)
        if self.total > 8 * LIMIT:
            raise InventoryFailure('inventory_output_limit')
        return decode(raw) if parsed else raw


def host_proof(read):
    if (os.environ.get('DOCKER_HOST', '') not in ('', 'unix:///var/run/docker.sock', 'unix:///run/docker.sock')
            or os.environ.get('DOCKER_CONTEXT', '') not in ('', 'default')
            or any(os.environ.get(key) for key in
                   ('DOCKER_TLS_VERIFY', 'DOCKER_CERT_PATH', 'DOCKER_API_VERSION', 'DOCKER_CONFIG'))):
        raise InventoryFailure('docker_endpoint_unsupported')
    endpoint = read(['docker', 'context', 'inspect', '--format', '{{json .Endpoints.docker.Host}}'])
    if endpoint not in ('unix:///var/run/docker.sock', 'unix:///run/docker.sock'):
        raise InventoryFailure('docker_endpoint_unsupported')
    try:
        proof = read([sys.executable, '-c', HOST_PROBE, endpoint[len('unix://'):]])
    except InventoryFailure:
        raise InventoryFailure('host_namespace_unverified') from None
    if (not isinstance(proof, dict) or set(proof) != {'local_rootful', 'same_namespace'}
            or any(value is not True for value in proof.values())):
        raise InventoryFailure('host_namespace_unverified')


def engine_info(read, require_pools):
    info = read(['docker', 'info', '--format', INFO_FORMAT])
    compose = read(['docker', 'compose', 'version', '--short'], parsed=False)
    try:
        compose = compose.decode('ascii').strip()
        if compose.startswith('v'):
            compose = compose[1:]
    except UnicodeError:
        raise InventoryFailure('docker_engine_unsupported') from None
    if (not isinstance(info, dict) or set(info) != {'version', 'os', 'security', 'pools'}
            or not isinstance(info['version'], str)
            or not re.fullmatch(r'\d+\.\d+\.\d+(?:[-+][A-Za-z0-9._-]{1,64})?', info['version'])
            or not re.fullmatch(r'\d+\.\d+\.\d+(?:[-+][A-Za-z0-9._-]{1,64})?', compose) or info['os'] != 'linux'
            or not isinstance(info['security'], list)
            or any(not isinstance(value, str) or 'rootless' in value for value in info['security'])):
        raise InventoryFailure('docker_engine_unsupported')
    result = {'docker_version': info['version'], 'compose_version': compose,
              'python_version': '.'.join(str(value) for value in sys.version_info[:3]),
              'pools': [], 'pool_source': 'unverified', 'pool_failure': ''}
    try:
        result['pools'], result['pool_source'] = effective_pools(info['pools'], info['version'])
    except InventoryFailure as error:
        if require_pools:
            raise
        result['pool_failure'] = error.code
    return result


def network_ipam(value):
    if (not isinstance(value, dict) or set(value) - {'Driver', 'Options', 'Config'}
            or value.get('Driver') not in ('default', 'null') or not isinstance(value.get('Config'), list)):
        raise InventoryFailure('network_ipam_unsupported')
    options = value.get('Options')
    if options is not None and (not isinstance(options, dict) or options):
        raise InventoryFailure('network_ipam_unsupported')
    result, excluded = [], set()
    for row in value['Config']:
        if not isinstance(row, dict) or set(row) - {'Subnet', 'Gateway', 'IPRange', 'AuxiliaryAddresses'}:
            raise InventoryFailure('network_ipam_unsupported')
        subnet = prefix(row.get('Subnet'), strict=True)
        if subnet.version == 4:
            excluded.add(str(subnet))
        for key in ('Gateway', 'IPRange'):
            if row.get(key):
                address = prefix(row[key])
                if address.version != subnet.version or not address.subnet_of(subnet):
                    raise InventoryFailure('network_ipam_unsupported')
        auxiliary = row.get('AuxiliaryAddresses') or {}
        if not isinstance(auxiliary, dict):
            raise InventoryFailure('network_ipam_unsupported')
        for address in auxiliary.values():
            if not prefix(address).subnet_of(subnet):
                raise InventoryFailure('network_ipam_unsupported')
        result.append(dict(row))
    return {'Driver': value['Driver'], 'Options': options or {}, 'Config': result}, excluded


def network_record(row):
    if not isinstance(row, dict):
        raise InventoryFailure('inventory_invalid')
    result = {key: row.get(key) for key in ('Id', 'Name', 'Driver', 'Scope', 'Internal', 'Attachable', 'EnableIPv6')}
    if (not isinstance(result['Id'], str) or not re.fullmatch(r'[a-f0-9]{64}', result['Id'])
            or not isinstance(result['Name'], str) or not 1 <= len(result['Name']) <= 256
            or result['Driver'] not in ('bridge', 'host', 'null', 'overlay', 'macvlan', 'ipvlan')
            or result['Scope'] not in ('local', 'swarm', 'global')
            or any(type(result[key]) is not bool for key in ('Internal', 'Attachable', 'EnableIPv6'))):
        raise InventoryFailure('inventory_invalid')
    for key in ('Labels', 'Options'):
        value = row.get(key)
        if value is not None and (not isinstance(value, dict) or any(
                not isinstance(k, str) or not isinstance(v, str) for k, v in value.items())):
            raise InventoryFailure('inventory_invalid')
        result[key] = value or {}
    containers = row.get('Containers')
    if not isinstance(containers, dict) or any(not isinstance(item, dict) for item in containers.values()):
        raise InventoryFailure('inventory_invalid')
    result['Containers'] = {key: {} for key in containers}
    ipam = row.get('IPAM')
    if result['Driver'] in ('host', 'null') and isinstance(ipam, dict) and ipam.get('Config', []) is None:
        ipam = dict(ipam, Config=[])
    result['IPAM'], excluded = network_ipam(ipam)
    return result, excluded


def network_ids(raw):
    try:
        ids = raw.decode('ascii').splitlines()
    except UnicodeError:
        raise InventoryFailure('inventory_invalid') from None
    if (not ids or len(ids) > 1024 or len(set(ids)) != len(ids)
            or any(not re.fullmatch(r'[a-f0-9]{64}', value) for value in ids)):
        raise InventoryFailure('inventory_invalid')
    return ids


def network_scan(read):
    listing = ['docker', 'network', 'ls', '--no-trunc', '--format', '{{.ID}}']
    ids = network_ids(read(listing, parsed=False))
    networks, excluded = [], set()
    for offset in range(0, len(ids), 32):
        batch = ids[offset:offset + 32]
        rows = read(['docker', 'network', 'inspect', *batch])
        if not isinstance(rows, list) or len(rows) != len(batch):
            raise InventoryFailure('inventory_invalid')
        observed = []
        for row in rows:
            network, prefixes = network_record(row)
            networks.append(network)
            observed.append(network['Id'])
            excluded.update(prefixes)
        if sorted(observed) != sorted(batch):
            raise InventoryFailure('inventory_invalid')
    if sorted(network_ids(read(listing, parsed=False))) != sorted(ids):
        raise InventoryFailure('network_inventory_changed')
    return networks, excluded


def host_networks(read):
    routes = read(['ip', '-j', '-4', 'route', 'show', 'table', 'all'])
    rules = read(['ip', '-j', '-4', 'rule', 'show'])
    addresses = read(['ip', '-j', '-4', 'address', 'show'])
    resolver = read(['cat', '/etc/resolv.conf'], parsed=False)
    validate_rules(rules)
    records = []
    excluded = route_exclusions(routes, records) | address_exclusions(addresses, records)
    dns, stub = resolver_exclusions(resolver, records=records)
    if stub:
        try:
            upstream = read(['resolvectl', 'dns'], parsed=False)
        except InventoryFailure:
            raise InventoryFailure('resolver_upstream_unknown') from None
        dns.update(resolver_exclusions(upstream, resolved=True, records=records)[0])
    return excluded | dns, sorted(records, key=lambda row: (row['subnet'], row['kind'], row['device']))


def collect(run, *, require_pools=True):
    """Return private fields only; callers must project their own closed public receipt."""
    read = Reader(run)
    host_proof(read)
    result = engine_info(read, False)
    try:
        if require_pools and result['pool_failure']:
            raise InventoryFailure(result['pool_failure'])
        before = host_networks(read)
        networks, occupied = network_scan(read)
        after = host_networks(read)
        if before != after:
            raise InventoryFailure('host_inventory_changed')
        host_proof(read)
        if engine_info(read, False) != result:
            raise InventoryFailure('host_inventory_changed')
    except InventoryFailure as error:
        error.details = dict(result, complete=False,
                             observed_at=datetime.datetime.now(datetime.timezone.utc).isoformat())
        raise
    result.update({'networks': networks, 'excluded': sorted(occupied | after[0]), 'host_excluded': sorted(after[0]),
                   'host_records': after[1],
                   'complete': True, 'observed_at': datetime.datetime.now(datetime.timezone.utc).isoformat()})
    return result
