26 lines
1.3 KiB
Python
26 lines
1.3 KiB
Python
import io,unittest,wave
|
|
from types import SimpleNamespace
|
|
from stt import validate_wav,read_upload,supported,REPO,REVISION,MODEL,STT
|
|
|
|
def audio():
|
|
out=io.BytesIO()
|
|
with wave.open(out,'wb') as w:w.setnchannels(1);w.setsampwidth(2);w.setframerate(16000);w.writeframes(b'\0\0'*1600)
|
|
return out.getvalue()
|
|
class STTTests(unittest.TestCase):
|
|
def test_wav(self):
|
|
validate_wav(audio())
|
|
for b in [b'',b'bad',audio()[:-2]]:
|
|
with self.assertRaises(ValueError):validate_wav(b)
|
|
def test_recipe(self):
|
|
self.assertTrue(supported(dict(repo=REPO,revision=REVISION,file=MODEL)))
|
|
self.assertFalse(supported(dict(repo=REPO,revision='other',file=MODEL)))
|
|
def test_multipart(self):
|
|
body=b'--x\r\nContent-Disposition: form-data; name="file"; filename="audio.wav"\r\n\r\n'+audio()+b'\r\n--x\r\nContent-Disposition: form-data; name="model"\r\n\r\nASR\r\n--x--\r\n'
|
|
h=SimpleNamespace(headers={'Content-Type':'multipart/form-data; boundary=x','Content-Length':str(len(body))},rfile=io.BytesIO(body),connection=SimpleNamespace(settimeout=lambda n:None))
|
|
fields,data=read_upload(h);self.assertEqual(fields['model'],'ASR');validate_wav(data)
|
|
h.headers['Transfer-Encoding']='chunked'
|
|
with self.assertRaises(ValueError):read_upload(h)
|
|
def test_missing_runtime(self):
|
|
s=STT(None,None,SimpleNamespace(status=lambda:{'active':None}))
|
|
with self.assertRaises(ValueError):s.build()
|