Files
Athena-Deck/test_separator.py
T

56 lines
4.2 KiB
Python

import io,json,tempfile,unittest,wave
from pathlib import Path
from unittest.mock import Mock,patch
from separator_runtime import SeparatorRuntime,MODELS
from separator import SeparatorTests
def audio(channels=2,duration=.1):
data=io.BytesIO()
with wave.open(data,'wb') as wav:
wav.setnchannels(channels);wav.setsampwidth(2);wav.setframerate(44100);wav.writeframes(b'\x00'*(int(44100*duration)*channels*2))
return data.getvalue()
class Tests(unittest.TestCase):
def test_marker_cannot_escape_and_failed_install_not_published(self):
with tempfile.TemporaryDirectory() as root:
Path(root,'active.json').write_text(json.dumps({'id':'../../outside'}));runtime=SeparatorRuntime(root)
self.assertFalse(runtime.status()['installed']);runtime.job={'id':'a'*32,'state':'running'}
with patch.object(runtime,'_command',side_effect=RuntimeError('synthetic failure')):runtime._run()
self.assertEqual(runtime.job['state'],'failed');self.assertFalse(runtime.status()['installed'])
def test_model_requires_configuration_and_unknown_download_rejected(self):
with tempfile.TemporaryDirectory() as root:
runtime=SeparatorRuntime(root);models=Path(root,'models');models.mkdir();name=next(iter(MODELS));(models/name).write_bytes(b'weights')
self.assertFalse(runtime.model_ready(name));(models/MODELS[name][1]).write_text('model: {}');(models/(name+'.ready')).write_text('{}');self.assertTrue(runtime.model_ready(name))
with self.assertRaises(ValueError):runtime.download('../other')
with self.assertRaises(ValueError):runtime.download([])
def test_stereo_accepted_and_long_or_truncated_wav_rejected(self):
with tempfile.TemporaryDirectory() as root:
runtime=Mock();runtime.status.return_value={'installed':True};runtime.model_ready.return_value=True
worker=SeparatorTests(root,runtime,Mock());model=next(iter(MODELS))
with patch('separator.threading.Thread') as thread:
self.assertEqual(worker.start(model,audio())['state'],'running');thread.return_value.start.assert_called_once()
worker.job=None
for data in (audio(duration=31),audio()[:-2],b'not a wav'):
with self.assertRaises(ValueError):worker.start(model,data)
def test_busy_gpu_fails_without_spawning_and_releases_lease(self):
with tempfile.TemporaryDirectory() as root:
scheduler=Mock();release=Mock();scheduler.image_reservation.return_value=release
worker=SeparatorTests(root,Mock(),scheduler);worker.job={'id':'a'*32,'state':'running','started_at':0}
with patch('separator.probe',return_value=[{'name':'RTX 3060','processes':1,'free_mib':11000}]),patch('separator.subprocess.Popen') as process:worker._run(next(iter(MODELS)),audio())
self.assertEqual(worker.job['state'],'failed');process.assert_not_called();release.assert_called_once()
def test_success_publishes_only_finished_outputs_and_removes_input(self):
with tempfile.TemporaryDirectory() as root:
scheduler=Mock();release=Mock();scheduler.image_reservation.return_value=release
runtime=Mock();runtime.paths.return_value=(Path('/fake/python'),Path('/fake'));runtime.root=Path(root,'runtime')
worker=SeparatorTests(root,runtime,scheduler);worker.job={'id':'a'*32,'state':'running','started_at':0}
directory=Path(root,'a'*32)
def launch(*args,**kwargs):
(directory/'vocals.wav').write_bytes(audio());(directory/'results.json').write_text('["vocals.wav"]')
process=Mock();process.poll.return_value=0;process.returncode=0;return process
with patch('separator.probe',return_value=[{'name':'RTX 3060','uuid':'GPU-test','processes':0,'free_mib':11000}]),patch('separator.subprocess.Popen',side_effect=launch):worker._run(next(iter(MODELS)),audio())
self.assertEqual(worker.job['state'],'complete');self.assertEqual(worker.audio('a'*32,0),audio());self.assertFalse((directory/'input.wav').exists());self.assertFalse((directory/'worker.tmp').exists());release.assert_called_once()
def test_result_traversal_rejected(self):
with tempfile.TemporaryDirectory() as root:
worker=SeparatorTests(root,Mock(),Mock());directory=Path(root,'a'*32);directory.mkdir();(directory/'complete').touch();(directory/'results.json').write_text('["../../secret.wav"]')
with self.assertRaises(ValueError):worker.audio('a'*32,0)
if __name__=='__main__':unittest.main()