254 lines
8.9 KiB
Python
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())
|