| from __future__ import annotations |
|
|
| import sys |
| import unittest |
| from copy import deepcopy |
| from pathlib import Path |
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) |
|
|
| from turn_detection.data import GroupingConfig, attach_group_ids, derive_group_keys |
|
|
|
|
| class LeakageGroupingTests(unittest.TestCase): |
| def test_transitive_audio_and_speaker_links(self) -> None: |
| raw = [ |
| {"id": "a", "dataset": "d", "speaker_id": "speaker-one", "audio_sha256": "x"}, |
| {"id": "b", "dataset": "d", "speaker_id": "speaker-two", "audio_sha256": "x"}, |
| {"id": "c", "dataset": "d", "speaker_id": "speaker-two", "audio_sha256": "y"}, |
| {"id": "d", "dataset": "d", "speaker_id": "speaker-three", "audio_sha256": "z"}, |
| ] |
| rows = [ |
| { |
| **record, |
| "group_keys": derive_group_keys(record, config=GroupingConfig(include_text=False)), |
| } |
| for record in raw |
| ] |
| attach_group_ids(rows) |
| self.assertEqual(rows[0]["group_id"], rows[1]["group_id"]) |
| self.assertEqual(rows[1]["group_id"], rows[2]["group_id"]) |
| self.assertNotEqual(rows[2]["group_id"], rows[3]["group_id"]) |
|
|
| def test_group_ids_are_order_independent(self) -> None: |
| rows = [ |
| {"record_id": "a", "group_keys": ["audio:x", "speaker:1"]}, |
| {"record_id": "b", "group_keys": ["speaker:1"]}, |
| {"record_id": "c", "group_keys": ["audio:z"]}, |
| ] |
| forward = attach_group_ids(deepcopy(rows)) |
| reverse = attach_group_ids(list(reversed(deepcopy(rows)))) |
| forward_ids = {row["record_id"]: row["group_id"] for row in forward} |
| reverse_ids = {row["record_id"]: row["group_id"] for row in reverse} |
| self.assertEqual(forward_ids, reverse_ids) |
|
|
| def test_identifiers_are_hashed_in_manifest_keys(self) -> None: |
| keys = derive_group_keys( |
| {"id": "record-1", "dataset": "source", "speaker_id": "private-speaker"} |
| ) |
| joined = " ".join(keys) |
| self.assertNotIn("private-speaker", joined) |
| self.assertTrue(any(key.startswith("speaker:") for key in keys)) |
|
|
| def test_repeated_nontrivial_text_links_prompts(self) -> None: |
| first = derive_group_keys( |
| {"id": "1", "dataset": "a", "spoken_text": "mera order kidhar hai please"} |
| ) |
| second = derive_group_keys( |
| {"id": "2", "dataset": "b", "spoken_text": "Mera order kidhar hai, please?"} |
| ) |
| self.assertTrue(set(key for key in first if key.startswith("text:")) & set(second)) |
|
|
| def test_empty_keys_get_unique_stable_fallbacks(self) -> None: |
| rows = [ |
| {"record_id": "a", "source_file": "x", "source_row": 0, "group_keys": []}, |
| {"record_id": "b", "source_file": "x", "source_row": 1, "group_keys": []}, |
| ] |
| attach_group_ids(rows) |
| self.assertNotEqual(rows[0]["group_id"], rows[1]["group_id"]) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|