| from __future__ import annotations |
|
|
| import sys |
| import tempfile |
| import unittest |
| from pathlib import Path |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) |
|
|
| from turn_detection.data import ManifestRecordResolver, iter_manifest_records |
|
|
| try: |
| import pyarrow as pa |
| import pyarrow.parquet as pq |
| except ImportError: |
| pa = None |
| pq = None |
|
|
|
|
| @unittest.skipIf(pa is None or pq is None, "pyarrow is an optional data dependency") |
| class ManifestResolverTests(unittest.TestCase): |
| def test_resolves_rows_and_audio_from_parquet_provenance(self) -> None: |
| with tempfile.TemporaryDirectory() as temporary: |
| path = Path(temporary) / "tiny.parquet" |
| table = pa.table( |
| { |
| "id": ["zero", "one", "two"], |
| "audio": [b"audio-zero", b"audio-one", b"audio-two"], |
| "endpoint_bool": [False, True, False], |
| } |
| ) |
| pq.write_table(table, path, row_group_size=2) |
| rows = [ |
| {"source_file": "tiny.parquet", "source_row": 1}, |
| {"source_file": "tiny.parquet", "source_row": 2}, |
| ] |
| resolver = ManifestRecordResolver(source_root=temporary, max_cached_row_groups=1) |
| self.assertEqual(resolver.resolve(rows[0], columns=("id",))["id"], "one") |
| self.assertEqual(resolver.resolve_audio(rows[1]), b"audio-two") |
| resolved = list( |
| iter_manifest_records( |
| rows, |
| source_root=temporary, |
| columns=("id", "endpoint_bool"), |
| ) |
| ) |
| self.assertEqual([record["id"] for record in resolved], ["one", "two"]) |
| self.assertEqual([record["endpoint_bool"] for record in resolved], [True, False]) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|