import unittest
from unittest.mock import patch
import torch
from benchmark import timed_call, run


class BenchmarkTests(unittest.TestCase):
    def test_synchronization_order(self):
        calls = []
        times = iter([10, 12])
        def sync(): calls.append('sync')
        def clock():
            calls.append('clock')
            return next(times)
        def generate():
            calls.append('generate')
            return 'output'
        self.assertEqual(timed_call(generate, sync, clock), ('output', 2))
        self.assertEqual(calls, ['sync', 'clock', 'generate', 'sync', 'clock'])

    def test_actual_cpu_generation_never_synchronizes_cuda(self):
        torch.set_num_threads(1)
        with patch.object(torch.cuda, 'synchronize', side_effect=AssertionError('CPU touched CUDA')):
            result = run('cpu', repeats=2, new_tokens=3)
        self.assertEqual(result['parameters'], 16864)
        self.assertEqual(result['prompt_tokens'], 4)
        self.assertEqual(len(result['samples']), 2)
        for sample in result['samples']:
            self.assertEqual(sample['new_tokens'], 3)
            self.assertGreater(sample['seconds'], 0)
            self.assertAlmostEqual(sample['tokens_per_second'] * sample['seconds'], 3)

    def test_validation(self):
        for args in [('cpu', 0, 3), ('cpu', 1, 61), ('cpu', 1, True), ('auto', 1, 3)]:
            with self.assertRaises(ValueError):
                run(*args)

    def test_no_cuda_fallback(self):
        with patch.object(torch.cuda, 'is_available', return_value=False):
            with self.assertRaisesRegex(ValueError, 'unavailable'):
                run('cuda:0')

    def test_zero_elapsed_rejected(self):
        with self.assertRaises(ValueError):
            timed_call(lambda: 1, lambda: None, lambda: 0)


if __name__ == '__main__':
    unittest.main()
