sesa-gpu / tests /test_session_isolation.py
John6666's picture
Upload 44 files
81ba775 verified
Raw
History Blame Contribute Delete
5.21 kB
import json
import threading
import time
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from src import downloads, jobs, preparation, service
def _prepared_job(job, token):
prepared = job.input_dir / "00_input.wav"
prepared.write_bytes(job.job_id.encode("ascii"))
config = {
"schema": preparation.PREPARATION_SCHEMA,
"reproducibility_schema": preparation.REPRODUCIBILITY_SCHEMA,
"status": "PREPARED",
"job_id": job.job_id,
"access_token": token,
"preparation_wall_seconds": 0.1,
"inputs": [{"prepared_path": "input/00_input.wav"}],
"input_total_seconds": 5.0,
"models": {
"selected_filenames": ["UVR-De-Reverb-aufr33-jarredou.pth"],
"package_filenames": ["UVR-De-Reverb-aufr33-jarredou.pth"],
"ensemble_algorithm": "avg_wave",
"custom_records": [],
},
"output": {
"format": "FLAC", "bitrate": "Auto", "sample_rate": 44100,
"normalization_threshold": 0.9, "amplification_threshold": 0.0,
"single_stem": "All stems", "bundle_layout": "Flat outputs",
},
"parameters": {"pitch_shift": 0, "chunk_duration": 0, "allow_cpu_fallback": True, "log_level": "INFO"},
"duration": {"seconds": 30},
"providers": {},
"outputs": [],
}
config["config_sha256"] = preparation._digest_for_private_config(config)
(job.config_dir / "preparation.json").write_text(json.dumps(config), encoding="utf-8")
(job.config_dir / "preflight.json").write_text("{}\n", encoding="utf-8")
(job.config_dir / ".prepared").write_text(config["config_sha256"] + "\n", encoding="utf-8")
(job.logs_dir / f"sesa_prepare_{job.job_id}.log").write_text("prepared\n", encoding="utf-8")
return preparation.prepared_state(job.job_id, token)
def test_overlapping_prepared_jobs_keep_outputs_and_reproducibility_isolated(monkeypatch, tmp_path):
monkeypatch.setattr(jobs, "JOB_ROOT", tmp_path / "jobs")
first = jobs.create_job()
second = jobs.create_job()
states = [_prepared_job(first, "token-a"), _prepared_job(second, "token-b")]
barrier = threading.Barrier(2)
class FakeSeparator:
def __init__(self, output_dir):
self.output_dir = Path(output_dir)
def load_model(self, value):
del value
barrier.wait(timeout=5)
def separate(self, values):
del values
output = self.output_dir / "00_input_(Vocals)_model.flac"
output.write_bytes(self.output_dir.parent.name.encode("ascii"))
return [str(output)]
monkeypatch.setattr(service, "create_separator", lambda **kwargs: FakeSeparator(kwargs["output_dir"]))
monkeypatch.setattr(service.torch.cuda, "is_available", lambda: False)
monkeypatch.setattr(service, "release_accelerators", lambda separator=None: None)
with ThreadPoolExecutor(max_workers=2) as pool:
results = list(pool.map(service.run_prepared_job, states))
assert all(result[0].startswith("### Completed") for result in results)
output_paths = [Path(result[3][0]) for result in results]
assert output_paths[0].parent != output_paths[1].parent
assert first.job_id in output_paths[0].name or first.job_id in output_paths[1].name
assert second.job_id in output_paths[0].name or second.job_id in output_paths[1].name
for result, expected_job in zip(results, (first, second)):
output = Path(result[3][0])
repro = json.loads(Path(result[7]).read_text(encoding="utf-8"))
assert output.parent == expected_job.output_dir
assert repro["job_id"] == expected_job.job_id
assert all(expected_job.job_id in record["name"] for record in repro["outputs"])
other = second.job_id if expected_job is first else first.job_id
assert other not in json.dumps(repro)
preparation.cleanup_prepared_state(states[0])
assert not first.root.exists()
assert second.root.exists()
def test_overlapping_package_prefetch_uses_one_shared_cache_download(monkeypatch, tmp_path):
model_root = tmp_path / "models"
download_root = tmp_path / "downloads"
monkeypatch.setattr(downloads, "MODEL_ROOT", model_root)
monkeypatch.setattr(downloads, "DOWNLOAD_ROOT", download_root)
monkeypatch.setattr(downloads.shutil, "which", lambda name: "/usr/bin/audio-separator")
calls = []
calls_lock = threading.Lock()
def fake_run(command, **kwargs):
del command, kwargs
with calls_lock:
calls.append(time.monotonic())
time.sleep(0.1)
model_root.mkdir(parents=True, exist_ok=True)
(model_root / "UVR_MDXNET_KARA_2.onnx").write_bytes(b"model")
monkeypatch.setattr(downloads.subprocess, "run", fake_run)
with ThreadPoolExecutor(max_workers=2) as pool:
records = list(pool.map(
lambda _: downloads.prefetch_package_models(["UVR_MDXNET_KARA_2.onnx"]),
range(2),
))
assert len(calls) == 1
assert all(result[0]["cache_after"] is True for result in records)
assert records[0][0]["files"][0]["sha256"] == records[1][0]["files"][0]["sha256"]