chatterbox-tts/tests/test_app.py

158 lines
4.7 KiB
Python

import os
import time
from pathlib import Path
import pytest
from fastapi.testclient import TestClient
import app as chatterbox_app
from app import ModelManager, app
client = TestClient(app)
def speech_payload(**overrides):
payload = {
"input": "Hello from Chatterbox.",
"voice": "reference.wav",
}
payload.update(overrides)
return payload
def configure_voices(monkeypatch, tmp_path):
voices_dir = tmp_path / "voices"
voices_dir.mkdir()
monkeypatch.setattr(chatterbox_app, "VOICES_DIR", voices_dir)
return voices_dir
def test_health_returns_ok_without_calling_model_loader(monkeypatch):
def fail_if_called(*args, **kwargs):
raise AssertionError("model loader must not be called")
monkeypatch.setattr(ModelManager, "load_model", fail_if_called)
response = client.get("/health")
assert response.status_code == 200
assert response.json() == {"status": "ok"}
def test_models_lists_chatterbox_without_calling_model_loader(monkeypatch):
def fail_if_called(*args, **kwargs):
raise AssertionError("model loader must not be called")
monkeypatch.setattr(ModelManager, "load_model", fail_if_called)
response = client.get("/v1/models")
assert response.status_code == 200
assert response.json()["data"][0]["id"] == "chatterbox"
def test_speech_rejects_voice_path_traversal():
response = client.post("/v1/audio/speech", json=speech_payload(voice="../secret.wav"))
assert response.status_code == 400
def test_speech_rejects_absolute_voice_path():
response = client.post(
"/v1/audio/speech",
json=speech_payload(voice=str(Path(chatterbox_app.VOICES_DIR) / "reference.wav")),
)
assert response.status_code == 400
def test_speech_rejects_missing_voice_file():
response = client.post(
"/v1/audio/speech",
json=speech_payload(voice="definitely-not-present-reference.wav"),
)
assert response.status_code == 400
def test_speech_rejects_input_over_2000_characters():
response = client.post("/v1/audio/speech", json=speech_payload(input="x" * 2001))
assert response.status_code == 422
def test_speech_rejects_unsupported_response_format():
response = client.post(
"/v1/audio/speech", json=speech_payload(response_format="mp3")
)
assert response.status_code == 400
def test_speech_rejects_symlink_escaping_voices_directory(monkeypatch, tmp_path):
voices_dir = configure_voices(monkeypatch, tmp_path)
outside_voice = tmp_path / "outside.wav"
outside_voice.write_bytes(b"not really wav")
try:
(voices_dir / "escape.wav").symlink_to(outside_voice)
except OSError as error:
pytest.skip(f"symlinks are unavailable on this OS: {error}")
response = client.post("/v1/audio/speech", json=speech_payload(voice="escape.wav"))
assert response.status_code == 400
def test_speech_rejects_directory_as_voice_target(monkeypatch, tmp_path):
voices_dir = configure_voices(monkeypatch, tmp_path)
(voices_dir / "voice-dir").mkdir()
response = client.post("/v1/audio/speech", json=speech_payload(voice="voice-dir"))
assert response.status_code == 400
def test_speech_rejects_unreadable_voice_target(monkeypatch, tmp_path):
if os.geteuid() == 0:
pytest.skip("root can read files regardless of their permission bits")
voices_dir = configure_voices(monkeypatch, tmp_path)
voice = voices_dir / "unreadable.wav"
voice.write_bytes(b"not really wav")
voice.chmod(0)
try:
response = client.post("/v1/audio/speech", json=speech_payload(voice="unreadable.wav"))
finally:
voice.chmod(0o600)
assert response.status_code == 400
def test_speech_accepts_exactly_2000_characters(monkeypatch, tmp_path):
voices_dir = configure_voices(monkeypatch, tmp_path)
(voices_dir / "reference.wav").write_bytes(b"not really wav")
monkeypatch.setattr(chatterbox_app.model_manager, "generate", lambda *args: b"wav")
response = client.post("/v1/audio/speech", json=speech_payload(input="x" * 2000))
assert response.status_code == 200
assert response.headers["content-type"] == "audio/wav"
assert response.content == b"wav"
def test_generation_failure_still_records_idle_unload_time():
class FailingModel:
def generate(self, *args, **kwargs):
raise RuntimeError("synthesis failed")
manager = ModelManager()
manager._model = FailingModel()
before_request = time.monotonic()
with pytest.raises(RuntimeError, match="synthesis failed"):
manager.generate("hello", Path("reference.wav"), 1.0)
try:
assert manager.last_request_time >= before_request
assert manager._unload_timer is not None
finally:
manager._cancel_unload_timer()