sesa-gpu / tests /test_dataset_samples.py
John6666's picture
Upload 44 files
81ba775 verified
Raw
History Blame Contribute Delete
3.83 kB
import zipfile
from pathlib import Path
from src import dataset_samples
class FakeApi:
def __init__(self, files):
self.files = files
self.calls = []
def list_repo_files(self, **kwargs):
self.calls.append(kwargs)
return list(self.files)
def test_dataset_catalog_has_licensed_default_and_fallbacks():
choices = dataset_samples.dataset_sample_choices()
assert choices
assert dataset_samples.DEFAULT_DATASET_SAMPLE_ID == "sonicsets-demo-mix"
assert any("CC BY-NC" in label for label, _ in choices)
assert len({value for _, value in choices}) == len(choices)
def test_mix_selector_prefers_full_ensemble_over_stems():
selected = dataset_samples._choose_audio_filename(
[
"Demo_01/Isolated_Vocals.wav",
"Demo_01/Full_Ensemble_Mix.wav",
"Demo_01/Bass.wav",
],
"ensemble-mix",
)
assert selected == "Demo_01/Full_Ensemble_Mix.wav"
def test_resolve_dataset_sample_downloads_only_selected_audio(monkeypatch, tmp_path):
api = FakeApi([
"README.md",
"Demo_01/Isolated_Vocals.wav",
"Demo_01/Full_Ensemble_Mix.wav",
])
calls = []
def fake_download(**kwargs):
calls.append(kwargs)
target = Path(kwargs["local_dir"]) / kwargs["filename"]
target.parent.mkdir(parents=True, exist_ok=True)
target.write_bytes(b"audio")
return str(target)
monkeypatch.setattr(dataset_samples, "DATASET_SAMPLE_ROOT", tmp_path)
result = dataset_samples.resolve_dataset_sample(
"sonicsets-demo-mix", api=api, download_fn=fake_download
)
assert result.repository_filename == "Demo_01/Full_Ensemble_Mix.wav"
assert result.license_id == "CC-BY-NC-4.0"
assert result.archive_extracted is False
assert len(calls) == 1
assert calls[0]["repo_type"] == "dataset"
def test_resolve_dataset_sample_can_extract_archive(monkeypatch, tmp_path):
api = FakeApi(["stems_evaluation.zip"])
def fake_download(**kwargs):
target = Path(kwargs["local_dir"]) / kwargs["filename"]
target.parent.mkdir(parents=True, exist_ok=True)
with zipfile.ZipFile(target, "w") as archive:
archive.writestr("Demo_01/Full_Ensemble_Mix.wav", b"mix")
archive.writestr("Demo_01/Isolated_Vocals.wav", b"vocal")
return str(target)
monkeypatch.setattr(dataset_samples, "DATASET_SAMPLE_ROOT", tmp_path)
result = dataset_samples.resolve_dataset_sample(
"sonicsets-demo-mix", api=api, download_fn=fake_download
)
assert result.archive_extracted is True
assert Path(result.local_path).name == "Full_Ensemble_Mix.wav"
assert Path(result.local_path).read_bytes() == b"mix"
def test_safe_extract_rejects_path_traversal(tmp_path):
archive_path = tmp_path / "unsafe.zip"
with zipfile.ZipFile(archive_path, "w") as archive:
archive.writestr("../escape.wav", b"bad")
try:
dataset_samples._safe_extract_zip(archive_path, tmp_path / "out")
except RuntimeError as exc:
assert "Unsafe path" in str(exc)
else:
raise AssertionError("unsafe archive path was accepted")
def test_resolve_rejects_oversized_download(monkeypatch, tmp_path):
target = tmp_path / "large.mp3"
target.write_bytes(b"12345")
def fake_download(**_kwargs):
return str(target)
monkeypatch.setattr(dataset_samples, "DATASET_SAMPLE_ROOT", tmp_path / "cache")
monkeypatch.setattr(dataset_samples, "MAX_DATASET_SAMPLE_FILE_BYTES", 4)
try:
dataset_samples.resolve_dataset_sample(
"legacy-symphony-01", download_fn=fake_download
)
except RuntimeError as exc:
assert "file-size limit" in str(exc)
else:
raise AssertionError("oversized dataset sample was accepted")