import unittest
from worker import add, vector
import torch


class WorkerTests(unittest.TestCase):
    def test_limits(self):
        self.assertEqual(len(vector([0] * 4096)), 4096)
        for value in [[0] * 4097, [], [True], [float('nan')], [float('inf')], [1e21]]:
            with self.assertRaises(ValueError):
                vector(value)

    def test_schema(self):
        for request in [[], {'a': [1]}, {'a': [1], 'b': [2], 'command': 'extra'}]:
            with self.assertRaises(ValueError):
                add(request, 'cpu')

    def test_no_silent_fallback(self):
        if not torch.cuda.is_available():
            with self.assertRaisesRegex(ValueError, 'no silent CPU fallback'):
                add({'a': [1], 'b': [2]}, 'cuda:0')


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