vignette/scripts/melotts-server.py
2026-08-09 18:22:03 +09:00

254 lines
8.9 KiB
Python

# -*- coding: utf-8 -*-
"""로컬 상주 MeloTTS 한국어 TTS 사이드카.
Higgs 서버와 같은 loopback 상주 방식이지만 권리 조건이 다르다. MeloTTS 는 MIT
라이선스라 상업·비상업 사용에 제약이 없고, 사전학습된 한국어 다화자 모델이라
실존 인물 음성 reference 를 전혀 쓰지 않는다. 그래서 Higgs 처럼 dev 전용 가드나
P1 프리셋 한정이 필요 없다.
엔드포인트:
GET /health → {"status":"ok","model":...,"language":"KR","license":"MIT",...}
POST /tts → {"text": "...", "speed": 1.0} 를 받아 WAV 바이트를 돌려준다
텍스트는 메모리에서만 다루고 디스크에 쓰지 않는다.
"""
from __future__ import annotations
import argparse
import array
import io
import json
import math
import sys
import wave
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, HTTPServer
from typing import Any, Iterable, Protocol
DEFAULT_HOST = "127.0.0.1"
DEFAULT_PORT = 9883
DEFAULT_LANGUAGE = "KR"
DEFAULT_SPEED = 1.0
MIN_SPEED = 0.5
MAX_SPEED = 2.0
MAX_TEXT_CHARS = 2_000
MAX_BODY_BYTES = 64 * 1024
MODEL_ID = "melotts-korean"
LICENSE_ID = "MIT"
class TtsError(RuntimeError):
"""합성 실패. 입력 텍스트는 로그·응답에 다시 싣지 않는다."""
def __init__(self, code: str, *, http_status: int = 422):
self.code = code
self.http_status = http_status
super().__init__(code)
@dataclass(frozen=True, slots=True)
class TtsRequest:
text: str
speed: float
class Synthesizer(Protocol):
sample_rate: int
def synthesize(self, text: str, *, speed: float) -> Iterable[float]: ...
def clamp_speed(value: Any) -> float:
"""말 속도를 안전 범위로 접는다. 숫자가 아니면 기본값."""
try:
speed = float(value)
except (TypeError, ValueError):
return DEFAULT_SPEED
if math.isnan(speed) or math.isinf(speed):
return DEFAULT_SPEED
return max(MIN_SPEED, min(MAX_SPEED, speed))
def parse_tts_request(payload: Any) -> TtsRequest:
"""요청 본문을 검증한다. 빈 텍스트와 과대 입력은 fail-closed."""
if not isinstance(payload, dict):
raise TtsError("invalid_body")
text = payload.get("text")
if not isinstance(text, str):
raise TtsError("text_required")
text = text.strip()
if not text:
raise TtsError("text_required")
if len(text) > MAX_TEXT_CHARS:
raise TtsError("text_too_long", http_status=413)
return TtsRequest(text=text, speed=clamp_speed(payload.get("speed", DEFAULT_SPEED)))
def encode_wav(samples: Iterable[float], sample_rate: int) -> bytes:
"""float(-1..1) 시퀀스를 16-bit mono WAV 로 만든다."""
if sample_rate <= 0:
raise TtsError("invalid_sample_rate", http_status=500)
pcm = array.array("h")
for sample in samples:
value = float(sample)
if math.isnan(value):
value = 0.0
clipped = max(-1.0, min(1.0, value))
pcm.append(int(round(clipped * 32767)))
if sys.byteorder != "little": # pragma: no cover - little-endian 개발 환경
pcm.byteswap()
buffer = io.BytesIO()
with wave.open(buffer, "wb") as handle:
handle.setnchannels(1)
handle.setsampwidth(2)
handle.setframerate(sample_rate)
handle.writeframes(pcm.tobytes())
return buffer.getvalue()
def health_payload(synthesizer: Synthesizer, *, speakers: Iterable[str]) -> dict[str, Any]:
return {
"status": "ok",
"provider": "melotts",
"model": MODEL_ID,
"language": DEFAULT_LANGUAGE,
"license": LICENSE_ID,
"reference_policy": "pretrained-multispeaker-no-external-reference",
"sample_rate": synthesizer.sample_rate,
"speakers": sorted(speakers),
}
class MeloSynthesizer:
"""상주 MeloTTS 모델 하나. 텍스트는 메모리에서만 다룬다."""
def __init__(self, *, device: str, language: str = DEFAULT_LANGUAGE) -> None:
try:
from melo.api import TTS
except ImportError as exc: # pragma: no cover - 런타임 환경 의존
raise TtsError("melotts_unavailable", http_status=500) from exc
self._model = TTS(language=language, device=device)
self.device = device
self.language = language
self.speaker_ids = dict(self._model.hps.data.spk2id)
if not self.speaker_ids: # pragma: no cover - 모델 무결성
raise TtsError("melotts_no_speaker", http_status=500)
self._speaker_id = next(iter(self.speaker_ids.values()))
self.sample_rate = int(self._model.hps.data.sampling_rate)
def synthesize(self, text: str, *, speed: float) -> Iterable[float]:
audio = self._model.tts_to_file(
text, self._speaker_id, None, speed=speed, quiet=True
)
return audio
def make_handler(
synthesizer: Synthesizer, *, speakers: Iterable[str]
) -> type[BaseHTTPRequestHandler]:
speaker_list = list(speakers)
class Handler(BaseHTTPRequestHandler):
server_version = "VignetteMeloTTS/1"
sys_version = ""
def log_message(self, format: str, *args: object) -> None:
return None
def _json(self, status: int, payload: dict[str, Any]) -> None:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
self.end_headers()
self.wfile.write(body)
def do_GET(self) -> None: # noqa: N802
if self.path != "/health":
self._json(404, {"detail": "not_found"})
return
self._json(200, health_payload(synthesizer, speakers=speaker_list))
def do_POST(self) -> None: # noqa: N802
if self.path != "/tts":
self._json(404, {"detail": "not_found"})
return
try:
length = int(self.headers.get("Content-Length", "0"))
except ValueError:
self._json(400, {"detail": "invalid_content_length"})
return
if length < 1 or length > MAX_BODY_BYTES:
self._json(413, {"detail": "request_size_rejected"})
return
try:
payload = json.loads(self.rfile.read(length))
except (UnicodeDecodeError, ValueError):
self._json(400, {"detail": "invalid_json"})
return
try:
request = parse_tts_request(payload)
audio = synthesizer.synthesize(request.text, speed=request.speed)
wav = encode_wav(audio, synthesizer.sample_rate)
except TtsError as exc:
self._json(exc.http_status, {"detail": exc.code})
return
except Exception:
self._json(500, {"detail": "synthesis_failed"})
return
self.send_response(200)
self.send_header("Content-Type", "audio/wav")
self.send_header("Content-Length", str(len(wav)))
self.send_header("Cache-Control", "no-store")
self.send_header("X-Vignette-TTS-Provider", "melotts")
self.send_header("X-Vignette-TTS-Model", MODEL_ID)
self.send_header("X-Vignette-TTS-License", LICENSE_ID)
self.end_headers()
self.wfile.write(wav)
return Handler
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--host", default=DEFAULT_HOST, choices=[DEFAULT_HOST, "localhost"])
parser.add_argument("--port", type=int, default=DEFAULT_PORT)
parser.add_argument("--device", default="auto")
parser.add_argument("--language", default=DEFAULT_LANGUAGE)
return parser
def main(argv: Iterable[str] | None = None) -> int: # pragma: no cover - CLI
args = build_parser().parse_args(list(argv) if argv is not None else None)
synthesizer = MeloSynthesizer(device=args.device, language=args.language)
print(
json.dumps(
{
"ready": True,
"host": args.host,
"port": args.port,
"model": MODEL_ID,
"license": LICENSE_ID,
"language": synthesizer.language,
"device": synthesizer.device,
"sample_rate": synthesizer.sample_rate,
"speakers": sorted(synthesizer.speaker_ids),
},
ensure_ascii=False,
separators=(",", ":"),
),
flush=True,
)
handler = make_handler(synthesizer, speakers=synthesizer.speaker_ids)
HTTPServer((args.host, args.port), handler).serve_forever()
return 0
if __name__ == "__main__": # pragma: no cover - CLI
raise SystemExit(main())