#!/usr/bin/env python3
"""Original, offline, single-speaker word-chunk -> WebVTT/JSON example.

This is a subtitle adapter, not an ASR model or an audio-quality benchmark.
Only Python's standard library is used. See README.md for the bounded contract.
"""
from __future__ import annotations

import argparse
from dataclasses import asdict, dataclass
from decimal import Decimal, ROUND_HALF_UP, localcontext, InvalidOperation
import json
from pathlib import Path
import re
import unicodedata

MAX_INPUT_BYTES = 2 * 1024 * 1024
MAX_WORDS = 10000
MAX_SECONDS = Decimal('86400')
PUNCTUATION_END = re.compile(r'''[.!?,;:]["'\u201d\u2019)\]]*$''')


class InputError(ValueError):
    """An input is outside this adapter's deliberately narrow contract."""


@dataclass(frozen=True)
class Config:
    line_codepoints: int = 40
    max_lines: int = 2
    pause_ms: int = 600
    min_duration_ms: int = 1000
    max_duration_ms: int = 6000
    max_cps: int = 20

    def __post_init__(self):
        bounds = {
            'line_codepoints': (1, 120), 'max_lines': (1, 2),
            'pause_ms': (1, 60000), 'min_duration_ms': (1, 86400000),
            'max_duration_ms': (1, 86400000), 'max_cps': (1, 1000),
        }
        for name, (low, high) in bounds.items():
            value = getattr(self, name)
            if type(value) is not int or not low <= value <= high:
                raise InputError(f'{name} must be an integer in [{low}, {high}]')
        if self.min_duration_ms > self.max_duration_ms:
            raise InputError('min_duration_ms must not exceed max_duration_ms')


def reject_constant(value):
    raise InputError(f'non-finite JSON number is not allowed: {value}')


def unique_object(pairs):
    result = {}
    for key, value in pairs:
        if key in result:
            raise InputError(f'duplicate JSON key is not allowed: {key}')
        result[key] = value
    return result


def loads_input(text: str):
    """Preserve JSON decimal literals; reject permissive JSON NaN/Infinity."""
    try:
        return json.loads(text, parse_float=Decimal, parse_constant=reject_constant,
                          object_pairs_hook=unique_object)
    except (json.JSONDecodeError, RecursionError, ValueError, InvalidOperation) as exc:
        if isinstance(exc, InputError):
            raise
        raise InputError(f'invalid JSON: {exc}') from exc


def read_input(path: Path):
    # A bounded read, rather than an unbounded read followed by a size check.
    with path.open('rb') as stream:
        data = stream.read(MAX_INPUT_BYTES + 1)
    if len(data) > MAX_INPUT_BYTES:
        raise InputError(f'input exceeds {MAX_INPUT_BYTES} bytes')
    try:
        return loads_input(data.decode('utf-8'))
    except UnicodeDecodeError as exc:
        raise InputError('input must be UTF-8') from exc


def seconds(value, label):
    if isinstance(value, bool) or not isinstance(value, (int, float, Decimal)):
        raise InputError(f'{label} must be a finite JSON number, not bool/null/string')
    number = value if isinstance(value, Decimal) else Decimal(str(value))
    if not number.is_finite() or not Decimal(0) <= number <= MAX_SECONDS:
        raise InputError(f'{label} must be finite and between 0 and 86400 seconds')
    return number


def milliseconds(number: Decimal) -> int:
    # Half-up, not Python round()'s ties-to-even. Never adds display padding.
    if number < Decimal('0.0005'):
        return 0
    with localcontext() as context:
        context.prec = max(32, len(number.as_tuple().digits) + 8)
        return int((number * 1000).to_integral_value(rounding=ROUND_HALF_UP))


def validate_words(document, config: Config):
    chunks = document.get('chunks') if isinstance(document, dict) else document
    if not isinstance(chunks, list):
        raise InputError('input must be a chunk list or an object with a chunks list')
    if len(chunks) > MAX_WORDS:
        raise InputError(f'input exceeds {MAX_WORDS} words')
    words = []
    previous_end = None
    for index, chunk in enumerate(chunks):
        label = f'chunks[{index}]'
        if not isinstance(chunk, dict):
            raise InputError(f'{label} must be an object')
        raw_text = chunk.get('text')
        if not isinstance(raw_text, str):
            raise InputError(f'{label}.text must be a string')
        text = raw_text.strip(' ')
        if not text or any(ch.isspace() for ch in text):
            raise InputError(f'{label}.text must contain exactly one space-separated word')
        if any(unicodedata.category(ch) in ('Cc', 'Cs') for ch in text):
            raise InputError(f'{label}.text contains a control character or surrogate')
        if len(text) > config.line_codepoints:
            raise InputError(f'{label}.text exceeds line_codepoints; never truncate a word')
        timing = chunk.get('timestamp')
        if not isinstance(timing, (list, tuple)) or len(timing) != 2:
            raise InputError(f'{label}.timestamp must have exactly two numbers')
        start = seconds(timing[0], f'{label}.timestamp[0]')
        end = seconds(timing[1], f'{label}.timestamp[1]')
        if end <= start:
            raise InputError(f'{label} end must be greater than start')
        # Check before quantization, so rounding cannot hide source overlap.
        if previous_end is not None and start < previous_end:
            raise InputError(f'{label} overlaps or precedes the previous word; single-speaker policy')
        start_ms, end_ms = milliseconds(start), milliseconds(end)
        if end_ms <= start_ms:
            raise InputError(f'{label} interval collapses after millisecond rounding')
        if words and start_ms < words[-1]['end_ms']:
            raise InputError(f'{label} overlaps after millisecond rounding')
        words.append({
            'source_index': index, 'text': text, 'source_text': raw_text,
            'source_start_seconds': str(start), 'source_end_seconds': str(end),
            'start_ms': start_ms, 'end_ms': end_ms,
        })
        previous_end = end
    return words


def wrap_words(words, config):
    """Greedy whole-word wrapping; None means line capacity is exhausted."""
    lines = []
    for word in words:
        text = word['text']
        if lines and len(lines[-1]) + 1 + len(text) <= config.line_codepoints:
            lines[-1] += ' ' + text
        elif len(lines) < config.max_lines:
            lines.append(text)
        else:
            return None
    return lines


def make_cue(words, ordinal, reason, config):
    text = ' '.join(word['text'] for word in words)
    start_ms, end_ms = words[0]['start_ms'], words[-1]['end_ms']
    duration = end_ms - start_ms
    cps = Decimal(len(text) * 1000) / Decimal(duration)
    issues = []
    if duration < config.min_duration_ms:
        issues.append('duration_below_example_minimum')
    if duration > config.max_duration_ms:
        issues.append('duration_above_example_maximum')
    if cps > Decimal(config.max_cps):
        issues.append('cps_above_example_maximum')
    return {
        'id': f'cue-{ordinal:04d}', 'start_ms': start_ms, 'end_ms': end_ms,
        'duration_ms': duration, 'text': text, 'lines': wrap_words(words, config),
        'source_word_indices': [word['source_index'] for word in words],
        'boundary_reason': reason, 'codepoint_count_including_spaces': len(text),
        'characters_per_second': float(cps.quantize(Decimal('0.001'), rounding=ROUND_HALF_UP)),
        'review_flags': issues,
    }


def format_clock(ms):
    if type(ms) is not int or ms < 0:
        raise InputError('clock requires nonnegative integer milliseconds')
    total_seconds, fraction = divmod(ms, 1000)
    total_minutes, sec = divmod(total_seconds, 60)
    hours, minute = divmod(total_minutes, 60)
    return f'{hours:02d}:{minute:02d}:{sec:02d}.{fraction:03d}'


def escape_vtt(text):
    return text.replace('&', '&amp;').replace('<', '&lt;').replace('>', '&gt;')


def render_vtt(cues):
    blocks = ['WEBVTT']
    for cue in cues:
        blocks.append(cue['id'] + '\n' + format_clock(cue['start_ms']) + ' --> ' +
                      format_clock(cue['end_ms']) + '\n' +
                      '\n'.join(escape_vtt(line) for line in cue['lines']))
    return '\n\n'.join(blocks) + '\n\n'


def convert(document, config=None):
    config = config or Config()
    words = validate_words(document, config)
    cues, pending = [], []

    def flush(reason):
        if pending:
            cues.append(make_cue(pending, len(cues) + 1, reason, config))
            pending.clear()

    for word in words:
        if pending and word['start_ms'] - pending[-1]['end_ms'] >= config.pause_ms:
            flush('pause')
        if pending and wrap_words(pending + [word], config) is None:
            flush('line_capacity')
        pending.append(word)
        if PUNCTUATION_END.search(word['text']):
            flush('punctuation')
    flush('end_of_input')
    manifest = {
        'schema_version': 1,
        'adapter_scope': 'single_speaker_space_separated_words',
        'evidence_kind': 'adapter_output_not_asr_validation',
        'timing_policy': 'seconds_to_integer_ms_decimal_half_up_no_padding',
        'cue_id_policy': 'ordinal_deterministic_for_same_input_and_config_not_durable_after_resegmentation',
        'manual_editorial_review_required': True,
        'config': asdict(config),
        'source_word_count': len(words), 'cue_count': len(cues),
        'words': words, 'cues': cues,
    }
    return render_vtt(cues), manifest


def manifest_json(manifest):
    return json.dumps(manifest, ensure_ascii=False, sort_keys=True, indent=2, allow_nan=False) + '\n'


def main(argv=None):
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('input', type=Path, help='UTF-8 JSON word chunks; no audio input')
    parser.add_argument('--vtt', type=Path, required=True)
    parser.add_argument('--manifest', type=Path, required=True)
    for name, value in asdict(Config()).items():
        parser.add_argument('--' + name.replace('_', '-'), type=int, default=value)
    args = parser.parse_args(argv)
    try:
        paths = [args.input.resolve(), args.vtt.resolve(), args.manifest.resolve()]
        if len(set(paths)) != 3 or any(
            first.exists() and second.exists() and first.samefile(second)
            for index, first in enumerate(paths) for second in paths[index + 1:]
        ):
            raise InputError('input, VTT, and manifest paths must be distinct')
        config = Config(**{key: getattr(args, key) for key in asdict(Config())})
        vtt, manifest = convert(read_input(args.input), config)
        # Validation completes before output files are opened. The two writes
        # are not a filesystem transaction: an I/O error can leave one output.
        args.vtt.write_text(vtt, encoding='utf-8', newline='\n')
        args.manifest.write_text(manifest_json(manifest), encoding='utf-8', newline='\n')
    except (InputError, OSError) as exc:
        parser.exit(2, f'error: {exc}\n')
    print(f'{manifest["source_word_count"]} words -> {manifest["cue_count"]} cues; manual review required')
    return 0


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