#!/usr/bin/env python3
"""Read one app's fixed Firebase refusal labels; raw host output never leaves this process."""
import argparse
import datetime
import json
import os
from pathlib import Path
import re
import selectors
import shlex
import signal
import subprocess
import time

STAGES = frozenset({'auth_signature', 'app_check_signature', 'app_scope', 'account_lookup', 'unclassified'})
REASONS = frozenset({
    'unclassified', 'transport_timeout', 'transport_failure', 'token_invalid',
    'app_scope_invalid', 'token_size_invalid', 'token_header_invalid', 'token_time_invalid',
    'authentication_time_invalid', 'app_audience_invalid', 'subject_invalid', 'invalid_key_id',
    'key_fetch_busy', 'key_refresh_limited', 'unknown_key', 'key_provider_unavailable',
    'key_response_limit', 'key_set_invalid', 'key_identity_invalid', 'key_algorithm_invalid',
    'key_strength_invalid', 'account_verification_refused', 'account_response_limit', 'account_missing',
    'account_disabled_or_mismatched', 'account_revocation_boundary_missing', 'account_revoked',
})
SCHEMA = 'gowalk-cicd/backend-identity-diagnostics.v1'
LABEL = re.compile(r'(?:WARNING:[A-Za-z_][A-Za-z0-9_.-]{0,127}:)?'
                   r'firebase_identity_refused stage=([a-z_]+) reason=([a-z_]+)')
LIMIT = 1024 * 1024
FAILURES = frozenset({
    'unclassified', 'invalid_scope', 'invalid_window', 'invalid_destination',
    'container_ambiguous_or_absent', 'container_scope_mismatch', 'command_timeout',
    'output_limit', 'command_failed', 'transport_failure', 'remote_invalid',
    'artifact_write_failed',
})


class ReaderFailure(Exception):
    def __init__(self, code):
        self.code = code if isinstance(code, str) and code in FAILURES else 'unclassified'
        super().__init__(self.code)


def transfer(selector, process, payload, offset, data):
    for key, _ in selector.select(0.1):
        if key.fileobj is process.stdin:
            offset += os.write(process.stdin.fileno(), payload[offset:offset + 4096])
            if offset == len(payload):
                selector.unregister(process.stdin)
                process.stdin.close()
        else:
            chunk = os.read(process.stdout.fileno(), 65536)
            if not chunk:
                selector.unregister(process.stdout)
            data.extend(chunk)
            if len(data) > LIMIT:
                raise ReaderFailure('output_limit')
    return offset


def bounded(argv, *, payload=None, seconds=30, with_status=False):
    """Bound output, runtime and owned process cleanup without printing stderr."""
    process = subprocess.Popen(argv, stdin=subprocess.PIPE if payload else subprocess.DEVNULL,
                               stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, start_new_session=True)
    deadline, data = time.monotonic() + seconds, bytearray()
    try:
        with selectors.DefaultSelector() as selector:
            selector.register(process.stdout, selectors.EVENT_READ)
            if payload:
                os.set_blocking(process.stdin.fileno(), False)
                selector.register(process.stdin, selectors.EVENT_WRITE)
            offset = 0
            while selector.get_map():
                if time.monotonic() >= deadline:
                    raise ReaderFailure('command_timeout')
                offset = transfer(selector, process, payload, offset, data)
        try:
            status = process.wait(timeout=max(0.01, deadline - time.monotonic()))
        except subprocess.TimeoutExpired:
            raise ReaderFailure('command_timeout') from None
        if with_status:
            return status, bytes(data)
        if status:
            raise ReaderFailure('command_failed')
        return bytes(data)
    finally:
        # Descendants can retain a pipe after the direct child exits. Kill only this owned group.
        try:
            os.killpg(process.pid, signal.SIGKILL)
        except ProcessLookupError:
            pass
        process.wait(timeout=5)
        process.stdout.close()
        if process.stdin and not process.stdin.closed:
            process.stdin.close()


def validate(app, service, since):
    if not re.fullmatch(r'[a-z0-9][a-z0-9-]{0,79}', app):
        raise ReaderFailure('invalid_scope')
    if not re.fullmatch(r'[a-z0-9][a-z0-9_-]{0,63}', service):
        raise ReaderFailure('invalid_scope')
    if not re.fullmatch(r'\d{4}-\d\d-\d\dT\d\d:\d\d:\d\dZ', since):
        raise ReaderFailure('invalid_window')
    try:
        start = datetime.datetime.strptime(since, '%Y-%m-%dT%H:%M:%SZ').replace(tzinfo=datetime.timezone.utc)
    except ValueError:
        raise ReaderFailure('invalid_window') from None
    age = (datetime.datetime.now(datetime.timezone.utc) - start).total_seconds()
    if age < 0 or age > 86400:
        raise ReaderFailure('invalid_window')


def records(raw):
    result = []
    for line in raw.decode('utf-8', errors='replace').splitlines():
        match = LABEL.fullmatch(line)
        if match and match[1] in STAGES and match[2] in REASONS:
            result.append({'stage': match[1], 'reason': match[2]})
    return result[-100:]


def remote(app, service, since):
    validate(app, service, since)
    ids = bounded(['docker', 'ps', '--no-trunc', '--filter', f'label=com.docker.compose.project={app}',
                   '--filter', f'label=com.docker.compose.service={service}', '--format', '{{.ID}}'], seconds=10).splitlines()
    if len(ids) != 1 or not re.fullmatch(rb'[a-f0-9]{64}', ids[0]):
        raise ReaderFailure('container_ambiguous_or_absent')
    container = ids[0].decode('ascii')
    # Recheck exact labels; Docker filter matching alone does not prove the requested scope.
    inspected = json.loads(bounded(['docker', 'inspect', '--format', '{{json .Config.Labels}}', container],
                                   seconds=10))
    if (inspected.get('com.docker.compose.project') != app
            or inspected.get('com.docker.compose.service') != service):
        raise ReaderFailure('container_scope_mismatch')
    # docker logs writes the container's stderr to its own stderr. Merge only inside the bounded reader.
    command = ['docker', 'logs', '--since', since, '--tail', '1000', container]
    raw = bounded(['sh', '-c', 'exec "$@" 2>&1', 'diagnostics', *command], seconds=10)
    return {'schema': SCHEMA, 'ok': True, 'failure': '', 'records': records(raw)}


def local(args):
    validate(args.app, args.service, args.since)
    if (not re.fullmatch(r'[A-Za-z0-9][A-Za-z0-9.-]{0,252}', args.host or '')
            or not re.fullmatch(r'[a-z_][a-z0-9_-]{0,31}', args.user)):
        raise ReaderFailure('invalid_destination')
    command = shlex.join(['python3', '-', '--remote', '--app', args.app,
                          '--service', args.service, '--since', args.since])
    try:
        status, raw = bounded(['ssh', '-i', str(Path.home() / '.ssh/backend_deploy'), '-o', 'IdentitiesOnly=yes',
                       '-o', 'BatchMode=yes', '-o', 'ConnectTimeout=10', '-o', 'ServerAliveInterval=5',
                       '-o', 'ServerAliveCountMax=2', f'{args.user}@{args.host}', command],
                      payload=Path(__file__).read_bytes(), seconds=45, with_status=True)
    except OSError:
        raise ReaderFailure('transport_failure') from None
    return decode_remote(status, raw)


def decode_remote(status, raw):
    if status not in (0, 1):
        raise ReaderFailure('transport_failure')
    try:
        result = json.loads(raw)
    except (ValueError, TypeError):
        raise ReaderFailure('remote_invalid') from None
    if (not isinstance(result, dict) or set(result) != {'schema', 'ok', 'failure', 'records'}
            or result['schema'] != SCHEMA or not isinstance(result['ok'], bool)):
        raise ReaderFailure('remote_invalid')
    if status == 1 and result['ok'] is False and result['records'] == []:
        raise ReaderFailure(result['failure'])
    if status or result['ok'] is not True or result['failure'] != '':
        raise ReaderFailure('remote_invalid')
    if not isinstance(result['records'], list) or len(result['records']) > 100:
        raise ReaderFailure('remote_invalid')
    for row in result['records']:
        if (not isinstance(row, dict) or set(row) != {'stage', 'reason'}
                or not isinstance(row['stage'], str) or row['stage'] not in STAGES
                or not isinstance(row['reason'], str) or row['reason'] not in REASONS):
            raise ReaderFailure('remote_invalid')
    return result


def retain_result(path, result):
    """Retain only the already validated public envelope, never captured host output."""
    try:
        target = Path(path)
        target.parent.mkdir(parents=True, exist_ok=True)
        with target.open('x', encoding='utf-8') as output:
            output.write(json.dumps(result, separators=(',', ':')) + '\n')
    except OSError:
        raise ReaderFailure('artifact_write_failed') from None


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--remote', action='store_true')
    parser.add_argument('--app', required=True)
    parser.add_argument('--service', default='api')
    parser.add_argument('--since', required=True)
    parser.add_argument('--host')
    parser.add_argument('--user', default='root')
    parser.add_argument('--output', help='Exclusive local file for the closed diagnostic envelope')
    args = parser.parse_args()
    status = 0
    try:
        if args.remote and args.output:
            raise ReaderFailure('invalid_destination')
        result = remote(args.app, args.service, args.since) if args.remote else local(args)
        result = decode_remote(0, json.dumps(result).encode())
    except Exception as error:
        # Only our closed vocabulary survives; provider and parser exceptions remain unclassified.
        code = error.code if isinstance(error, ReaderFailure) else 'unclassified'
        result = {'schema': SCHEMA, 'ok': False, 'failure': code, 'records': []}
        status = 1
    if args.output and not args.remote:
        try:
            retain_result(args.output, result)
        except ReaderFailure as error:
            result = {'schema': SCHEMA, 'ok': False, 'failure': error.code, 'records': []}
            status = 1
    print(json.dumps(result, separators=(',', ':')))
    return status


if __name__ == '__main__':
    raise SystemExit(main())
