"""Synthetic contract tests; no inference, audio, network, or game engine."""
from decimal import Decimal
import json
import random
from pathlib import Path
import subprocess
import sys
import tempfile
import unittest

from subtitle_adapter import (Config, InputError, MAX_INPUT_BYTES, MAX_WORDS,
                              convert, format_clock, loads_input, manifest_json,
                              read_input)

BASE = Path(__file__).resolve().parent


def word(text='Hello', start=0, end=1):
    return {'text': text, 'timestamp': [start, end]}


class AdapterTests(unittest.TestCase):
    def reject(self, data, match=None, config=None):
        with self.assertRaisesRegex(InputError, match or '.'):
            convert(data, config)

    def test_empty_input(self):
        vtt, manifest = convert([])
        self.assertEqual(vtt, 'WEBVTT\n\n')
        self.assertEqual(manifest['cues'], [])

    def test_hf_object_envelope(self):
        self.assertEqual(convert({'text': 'ignored', 'chunks': [word()]}), convert([word()]))

    def test_top_level_types(self):
        for data in (None, 1, True, 'chunks', {}, {'chunks': None}, {'chunks': {}}):
            with self.subTest(data=data):
                self.reject(data)

    def test_chunk_must_be_object(self):
        for chunk in (None, True, [], 9, 'x'):
            with self.subTest(chunk=chunk):
                self.reject([chunk])

    def test_missing_fields(self):
        for chunk in ({}, {'text': 'ok'}, {'timestamp': [0, 1]}):
            with self.subTest(chunk=chunk):
                self.reject([chunk])

    def test_text_must_be_string(self):
        for text in (None, True, 1, [], {}):
            with self.subTest(text=text):
                self.reject([word(text)])

    def test_empty_text_rejected(self):
        for text in ('', ' ', '   '):
            with self.subTest(text=text):
                self.reject([word(text)])

    def test_one_word_only_not_segment_input(self):
        for text in ('two words', 'two\twords', 'a\nb', 'two\u00a0words', '\tword'):
            with self.subTest(text=text):
                self.reject([word(text)])

    def test_ascii_boundary_spaces_preserved_in_source(self):
        _, result = convert([word('  Hello  ')])
        self.assertEqual(result['cues'][0]['text'], 'Hello')
        self.assertEqual(result['words'][0]['source_text'], '  Hello  ')

    def test_control_and_surrogate_rejected(self):
        for text in ('a\x00b', 'a\x7fb', 'a\ud800b'):
            with self.subTest(text=repr(text)):
                self.reject([word(text)])

    def test_timing_shape(self):
        for timing in (None, [], [0], [0, 1, 2], '01', {'start': 0, 'end': 1}):
            with self.subTest(timing=timing):
                self.reject([{'text': 'ok', 'timestamp': timing}])

    def test_bool_time_rejected(self):
        for start, end in ((False, 1), (0, True)):
            self.reject([word(start=start, end=end)])

    def test_null_and_numeric_string_times_rejected(self):
        for start, end in ((None, 1), (0, None), ('0', 1), (0, '1')):
            self.reject([word(start=start, end=end)])

    def test_nonfinite_times_rejected(self):
        for value in (float('nan'), float('inf'), -float('inf'), Decimal('NaN'), Decimal('Infinity')):
            with self.subTest(value=value):
                self.reject([word(start=value)])
                self.reject([word(end=value)])

    def test_negative_time_rejected(self):
        self.reject([word(start=-0.001)])
        self.reject([word(start=-2, end=-1)])

    def test_nonincreasing_interval_rejected(self):
        self.reject([word(start=1, end=1)])
        self.reject([word(start=2, end=1)])

    def test_unsorted_words_rejected(self):
        self.reject([word(start=2, end=3), word(start=0, end=1)])

    def test_single_speaker_overlap_rejected(self):
        self.reject([word(start=0, end=1), word(start=0.9, end=2)], 'single-speaker policy')

    def test_touching_boundaries_accepted(self):
        _, result = convert([word(start=0, end=1), word(start=1, end=2)])
        self.assertEqual(result['source_word_count'], 2)

    def test_submillisecond_overlap_not_hidden(self):
        self.reject([word(end=Decimal('1.0004')), word(start=Decimal('1.0003'), end=2)], 'single-speaker policy')

    def test_rounding_half_up(self):
        _, result = convert([word(start=Decimal('0.0005'), end=Decimal('1.0005'))])
        self.assertEqual((result['words'][0]['start_ms'], result['words'][0]['end_ms']), (1, 1001))

    def test_rounding_does_not_create_extra_millisecond(self):
        data = loads_input('[{"text":"ok","timestamp":[0.000499999999999999999999999999999999999,1.000499999999999999999999999999999999999]}]')
        _, result = convert(data)
        self.assertEqual((result['words'][0]['start_ms'], result['words'][0]['end_ms']), (0, 1000))

    def test_quantization_collapse_rejected(self):
        self.reject([word(start=Decimal('0.0001'), end=Decimal('0.0004'))], 'collapses')
        self.reject([word(start=Decimal('1.0005'), end=Decimal('1.0014'))], 'collapses')

    def test_exact_milliseconds_preserved(self):
        _, result = convert(loads_input('[{"text":"ok","timestamp":[1.234,2.345]}]'))
        self.assertEqual((result['cues'][0]['start_ms'], result['cues'][0]['end_ms']), (1234, 2345))
        self.assertEqual(result['words'][0]['source_start_seconds'], '1.234')

    def test_original_decimal_seconds_retained(self):
        _, result = convert(loads_input('[{"text":"ok","timestamp":[0.12345,0.98765]}]'))
        self.assertEqual(result['words'][0]['source_start_seconds'], '0.12345')
        self.assertEqual(result['words'][0]['source_end_seconds'], '0.98765')

    def test_pause_at_threshold_splits(self):
        _, result = convert([word(end=1), word('again', 1.6, 2.6)])
        self.assertEqual(result['cue_count'], 2)
        self.assertEqual(result['cues'][0]['boundary_reason'], 'pause')

    def test_pause_below_threshold_does_not_split(self):
        _, result = convert([word(end=1), word('again', 1.599, 2.6)])
        self.assertEqual(result['cue_count'], 1)

    def test_punctuation_and_closing_quote_split(self):
        for text in ('Wait.', 'Wait,', 'Wait?', 'Wait!', 'Wait;', 'Wait:', 'Wait!\u201d'):
            with self.subTest(text=text):
                _, result = convert([word(text), word('next', 1, 2)])
                self.assertEqual(result['cue_count'], 2)
                self.assertEqual(result['cues'][0]['boundary_reason'], 'punctuation')

    def test_capacity_splits_without_dropping_words(self):
        data = [word(text, i, i + 1) for i, text in enumerate(['alpha', 'bravo', 'charlie', 'delta', 'echo'])]
        _, result = convert(data, Config(line_codepoints=11, max_lines=2))
        self.assertEqual(result['cue_count'], 2)
        self.assertEqual(result['cues'][0]['lines'], ['alpha bravo', 'charlie'])
        self.assertEqual(result['cues'][0]['boundary_reason'], 'line_capacity')
        self.assertEqual([i for cue in result['cues'] for i in cue['source_word_indices']], list(range(5)))
        self.assertEqual(' '.join(c['text'] for c in result['cues']), 'alpha bravo charlie delta echo')

    def test_exact_line_limit_accepted(self):
        _, result = convert([word('x' * 40)])
        self.assertEqual(len(result['cues'][0]['lines'][0]), 40)

    def test_oversize_word_rejected_never_truncated(self):
        self.reject([word('x' * 41)], 'never truncate')

    def test_unicode_is_codepoints_not_graphemes(self):
        # e + combining accent = 2 codepoints; this is not a font-width estimate.
        _, result = convert([word('e\u0301')], Config(line_codepoints=2))
        self.assertEqual(result['cues'][0]['codepoint_count_including_spaces'], 2)
        self.reject([word('e\u0301')], 'never truncate', Config(line_codepoints=1))
        # An emoji sequence can be more than one codepoint too.
        _, result = convert([word('\U0001f469\u200d\U0001f4bb')])
        self.assertEqual(result['cues'][0]['codepoint_count_including_spaces'], 3)

    def test_vtt_literal_escaping(self):
        vtt, result = convert([word('<map>&-->')])
        self.assertIn('&lt;map&gt;&amp;--&gt;', vtt)
        self.assertEqual(result['cues'][0]['text'], '<map>&-->')
        self.assertEqual(vtt.count('-->'), 1)  # Only the timing separator remains.

    def test_review_thresholds_flag_without_padding(self):
        _, result = convert([word('Urgent!', 0, 0.1)])
        cue = result['cues'][0]
        self.assertEqual(cue['review_flags'], ['duration_below_example_minimum', 'cps_above_example_maximum'])
        self.assertEqual((cue['start_ms'], cue['end_ms']), (0, 100))
        _, result = convert([word('Slow', 0, 8)])
        self.assertEqual(result['cues'][0]['review_flags'], ['duration_above_example_maximum'])

    def test_custom_thresholds_and_inclusive_boundaries(self):
        _, result = convert([word('abcd', 0, 1)], Config(min_duration_ms=1000, max_duration_ms=1000, max_cps=4))
        self.assertEqual(result['cues'][0]['review_flags'], [])
        _, result = convert([word('abcd', 0, 1)], Config(min_duration_ms=500, max_duration_ms=1000, max_cps=3))
        self.assertEqual(result['cues'][0]['review_flags'], ['cps_above_example_maximum'])

    def test_integer_clock_beyond_one_hour(self):
        self.assertEqual(format_clock(3723004), '01:02:03.004')
        self.assertEqual(format_clock(360000001), '100:00:00.001')
        vtt, _ = convert([word('Late.', 3723.004, 3724.005)])
        self.assertIn('01:02:03.004 --> 01:02:04.005', vtt)

    def test_clock_rejects_floats_and_negative(self):
        for value in (1.0, True, -1):
            with self.assertRaises(InputError):
                format_clock(value)

    def test_deterministic_serialization_and_ordinal_ids(self):
        data = [word('Go.'), word('Wait.', 1, 2)]
        first = convert(data)
        second = convert(data)
        self.assertEqual(first[0], second[0])
        self.assertEqual(manifest_json(first[1]), manifest_json(second[1]))
        self.assertEqual([cue['id'] for cue in first[1]['cues']], ['cue-0001', 'cue-0002'])
        self.assertIn('not_durable', first[1]['cue_id_policy'])

    def test_invalid_json_and_duplicate_keys(self):
        for text in ('{bad}', '[1,]', '1e9999999999999999999999999999999999', '{"chunks": [], "chunks": []}', '[{"text":"ok","text":"bad"}]'):
            with self.subTest(text=text), self.assertRaises(InputError):
                loads_input(text)

    def test_nonstandard_json_constants_rejected(self):
        for value in ('NaN', 'Infinity', '-Infinity'):
            with self.assertRaises(InputError):
                loads_input('[{"text":"ok","timestamp":[0,' + value + ']}]')

    def test_config_validation(self):
        for kwargs in ({'max_lines': 3}, {'max_lines': False}, {'line_codepoints': 0},
                       {'pause_ms': 0}, {'min_duration_ms': 7000, 'max_duration_ms': 6000},
                       {'max_cps': float('inf')}):
            with self.subTest(kwargs=kwargs), self.assertRaises(InputError):
                Config(**kwargs)

    def test_word_count_bound(self):
        self.reject([word()] * (MAX_WORDS + 1), 'exceeds')

    def test_timestamp_bound(self):
        self.reject([word(start=86400, end=86401)], '86400')
        _, result = convert([word(start=86399, end=86400)])
        self.assertEqual(result['cues'][0]['end_ms'], 86400000)

    def test_bounded_input_read_and_utf8(self):
        with tempfile.TemporaryDirectory(dir=BASE / 'results') as directory:
            path = Path(directory) / 'input.json'
            path.write_bytes(b' ' * (MAX_INPUT_BYTES + 1))
            with self.assertRaisesRegex(InputError, 'exceeds'):
                read_input(path)
            path.write_bytes(b'\xff')
            with self.assertRaisesRegex(InputError, 'UTF-8'):
                read_input(path)

    def test_cli_success(self):
        with tempfile.TemporaryDirectory(dir=BASE / 'results') as directory:
            root = Path(directory)
            source = root / 'in.json'
            source.write_text(json.dumps([word('Hello.')]), encoding='utf-8')
            result = subprocess.run([sys.executable, str(BASE / 'subtitle_adapter.py'), str(source),
                '--vtt', str(root / 'out.vtt'), '--manifest', str(root / 'out.json')], capture_output=True, text=True)
            self.assertEqual(result.returncode, 0, result.stderr)
            self.assertTrue((root / 'out.vtt').read_bytes().startswith(b'WEBVTT\n\n'))
            self.assertEqual(json.loads((root / 'out.json').read_text())['cue_count'], 1)

    def test_cli_validation_failure_writes_no_outputs(self):
        with tempfile.TemporaryDirectory(dir=BASE / 'results') as directory:
            root = Path(directory)
            source = root / 'in.json'
            source.write_text(json.dumps([word(end=None)]), encoding='utf-8')
            result = subprocess.run([sys.executable, str(BASE / 'subtitle_adapter.py'), str(source),
                '--vtt', str(root / 'out.vtt'), '--manifest', str(root / 'out.json')], capture_output=True, text=True)
            self.assertEqual(result.returncode, 2)
            self.assertFalse((root / 'out.vtt').exists())
            self.assertFalse((root / 'out.json').exists())

    def test_cli_rejects_same_paths(self):
        with tempfile.TemporaryDirectory(dir=BASE / 'results') as directory:
            root = Path(directory)
            source = root / 'in.json'
            original = json.dumps([word()])
            source.write_text(original, encoding='utf-8')
            result = subprocess.run([sys.executable, str(BASE / 'subtitle_adapter.py'), str(source),
                '--vtt', str(source), '--manifest', str(root / 'out.json')], capture_output=True, text=True)
            self.assertEqual(result.returncode, 2)
            self.assertEqual(source.read_text(), original)

    def test_exact_fixture_vtt_bytes(self):
        vtt, _ = convert(read_input(BASE / 'fixtures' / 'exact_milliseconds.json'))
        self.assertEqual(vtt, 'WEBVTT\n\n'
            'cue-0001\n00:00:00.001 --> 00:00:01.001\nReady.\n\n'
            'cue-0002\n00:00:01.234 --> 00:00:02.345\nSteady.\n\n'
            'cue-0003\n01:02:03.004 --> 01:02:04.005\nLater.\n\n')

    def test_seeded_synthetic_schedules_preserve_invariants(self):
        rng = random.Random(104)
        for _ in range(100):
            data, now = [], 0
            for index in range(rng.randint(1, 50)):
                start = now + rng.randint(0, 1000)
                now = start + rng.randint(1, 1200)
                text = 'word' + str(index) + rng.choice(['', '', ',', '.', '!'])
                data.append(word(text, Decimal(start) / 1000, Decimal(now) / 1000))
            _, result = convert(data)
            self.assertEqual([i for c in result['cues'] for i in c['source_word_indices']], list(range(len(data))))
            previous_end = 0
            for cue in result['cues']:
                self.assertGreater(cue['end_ms'], cue['start_ms'])
                self.assertGreaterEqual(cue['start_ms'], previous_end)
                self.assertLessEqual(len(cue['lines']), 2)
                self.assertTrue(all(len(line) <= 40 for line in cue['lines']))
                previous_end = cue['end_ms']

    def test_all_fixtures_preserve_every_word_and_boundary(self):
        for fixture in sorted((BASE / 'fixtures').glob('*.json')):
            with self.subTest(fixture=fixture.name):
                vtt, result = convert(read_input(fixture))
                cues, words = result['cues'], result['words']
                self.assertEqual([i for cue in cues for i in cue['source_word_indices']], list(range(len(words))))
                self.assertEqual(' '.join(cue['text'] for cue in cues), ' '.join(word['text'] for word in words))
                for cue in cues:
                    self.assertLess(cue['start_ms'], cue['end_ms'])
                    self.assertEqual(cue['start_ms'], words[cue['source_word_indices'][0]]['start_ms'])
                    self.assertEqual(cue['end_ms'], words[cue['source_word_indices'][-1]]['end_ms'])
                    self.assertLessEqual(len(cue['lines']), 2)
                    self.assertTrue(all(len(line) <= 40 for line in cue['lines']))
                self.assertTrue(vtt.endswith('\n\n'))


if __name__ == '__main__':
    (BASE / 'results').mkdir(exist_ok=True)
    unittest.main(verbosity=2)
