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
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1#!/usr/bin/env python3
3import os
4import pytest
5import binascii
7import torch
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
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)
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
27 tensor_binary = base64_to_binary(tensor_base64)
28 assert isinstance(tensor_binary, bytes)
30 tensor_base64 = binary_to_base64(tensor_binary)
31 assert isinstance(tensor_base64, str)
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
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)
44 with open("test_tensor.pt", "wb") as file:
45 file.write(tensor_binary)
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
52 os.remove("test_tensor.pt")
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")
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]