Coverage for tests/test_tensor_utils.py: 100%

50 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-09 04:47 +0000

1#!/usr/bin/env python3 

2 

3import os 

4import pytest 

5import binascii 

6 

7import torch 

8 

9from media_utils import tensor_to_base64 

10from file_utils import binary_to_base64 

11from file_utils import base64_to_binary 

12from media_utils import base64_to_tensor 

13from media_utils import bytes_to_tensor 

14from media_utils import get_tensor_file_info 

15 

16 

17def test_base64() -> None: 

18 """Test tensor to base64 and back conversion.""" 

19 tensor_data = torch.rand(3, 4, 5) 

20 tensor_base64 = tensor_to_base64(tensor_data) 

21 assert isinstance(tensor_base64, str) 

22 

23 tensor_data_2 = base64_to_tensor(tensor_base64) 

24 assert isinstance(tensor_data_2, torch.Tensor) 

25 assert tensor_data.shape == tensor_data_2.shape 

26 

27 tensor_binary = base64_to_binary(tensor_base64) 

28 assert isinstance(tensor_binary, bytes) 

29 

30 tensor_base64 = binary_to_base64(tensor_binary) 

31 assert isinstance(tensor_base64, str) 

32 

33 tensor_data_3 = base64_to_tensor(tensor_base64) 

34 assert isinstance(tensor_data_3, torch.Tensor) 

35 assert tensor_data.shape == tensor_data_3.shape 

36 

37 

38def test_tensor_file() -> None: 

39 """Test saving tensor to file and getting its info.""" 

40 tensor_data = torch.rand(3, 4, 5) 

41 tensor_base64 = tensor_to_base64(tensor_data) 

42 tensor_binary = base64_to_binary(tensor_base64) 

43 

44 with open("test_tensor.pt", "wb") as file: 

45 file.write(tensor_binary) 

46 

47 tensor_info = get_tensor_file_info("test_tensor.pt") 

48 assert tensor_info["dtype"].startswith("torch.float") 

49 assert tensor_info["shape"] == "torch.Size([3, 4, 5])" 

50 assert tensor_info["numel"] == 60 

51 

52 os.remove("test_tensor.pt") 

53 

54 with pytest.raises(TypeError): 

55 get_tensor_file_info(None) # type: ignore[arg-type] 

56 with pytest.raises(FileNotFoundError): 

57 get_tensor_file_info("nonexisting.pt") 

58 

59 

60def test_base64_invalid() -> None: 

61 """Test invalid inputs for base64 and tensor functions.""" 

62 with pytest.raises(TypeError): 

63 base64_to_binary(b"12345") # type: ignore[arg-type] 

64 with pytest.raises(TypeError): 

65 base64_to_tensor(12345) # type: ignore[arg-type] 

66 with pytest.raises(TypeError): 

67 tensor_to_base64("12345") # type: ignore[arg-type] 

68 with pytest.raises(binascii.Error): 

69 base64_to_tensor("NOTBASE64") 

70 with pytest.raises(TypeError): 

71 bytes_to_tensor("12345") # type: ignore[arg-type]