#!/usr/bin/env python3 from __future__ import annotations import argparse import json import os import signal import sys import time import traceback import uuid from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path from threading import Lock from typing import Any PROJECT_DIR = Path(__file__).resolve().parents[1] OUTPUT_DIR = Path(os.environ.get("KSAY_OUTPUT_DIR", PROJECT_DIR / "generated-audio")) DEFAULT_HOST = os.environ.get("KSAY_HOST", "127.0.0.1") DEFAULT_PORT = int(os.environ.get("KSAY_PORT", "7332")) DEFAULT_MODEL = os.environ.get("KSAY_MODEL", "mlx-community/Kokoro-82M-8bit") DEFAULT_VOICE = os.environ.get("KSAY_VOICE", "af_heart") DEFAULT_LANG_CODE = os.environ.get("KSAY_LANG_CODE", "a") MAX_TEXT_CHARS = int(os.environ.get("KSAY_MAX_TEXT_CHARS", "8000")) class KokoroEngine: def __init__(self, model_name: str): self.model_name = model_name self.model = None self.loaded_at = None self.load_seconds = None self.lock = Lock() def load(self) -> None: started = time.perf_counter() from mlx_audio.tts.utils import load_model self.model = load_model(self.model_name) self.loaded_at = time.time() self.load_seconds = time.perf_counter() - started def synthesize( self, *, text: str, voice: str, speed: float, lang_code: str, output: str | None, ) -> dict[str, Any]: if self.model is None: raise RuntimeError("Kokoro model is not loaded.") clean_text = text.strip() if not clean_text: raise ValueError("text is required.") clean_text = clean_text[:MAX_TEXT_CHARS] output_path = resolve_output_path(output) started = time.perf_counter() with self.lock: audio_chunks = [] sample_rate = None for result in self.model.generate( text=clean_text, voice=voice, speed=speed, lang_code=lang_code, ): audio_chunks.append(result.audio) sample_rate = result.sample_rate if not audio_chunks: raise RuntimeError("Kokoro did not return audio.") import mlx.core as mx import numpy as np from mlx_audio.audio_io import write as audio_write audio = ( mx.concatenate(audio_chunks, axis=0) if len(audio_chunks) > 1 else audio_chunks[0] ) audio_write(str(output_path), np.array(audio), sample_rate, format="wav") elapsed = time.perf_counter() - started return { "ok": True, "filePath": str(output_path), "model": self.model_name, "voice": voice, "speed": speed, "langCode": lang_code, "sampleRate": sample_rate, "segments": len(audio_chunks), "seconds": round(elapsed, 3), "characters": len(clean_text), } def resolve_output_path(output: str | None) -> Path: if output: path = Path(output).expanduser() if path.suffix.lower() != ".wav": path = path.with_suffix(".wav") if not path.is_absolute(): path = (Path.cwd() / path).resolve() else: OUTPUT_DIR.mkdir(parents=True, exist_ok=True) path = OUTPUT_DIR / f"ksay-{int(time.time() * 1000)}-{uuid.uuid4().hex[:6]}.wav" path.parent.mkdir(parents=True, exist_ok=True) return path def parse_json_body(handler: BaseHTTPRequestHandler) -> dict[str, Any]: length = int(handler.headers.get("content-length", "0")) if length <= 0: return {} body = handler.rfile.read(length) return json.loads(body.decode("utf-8")) def make_handler(engine: KokoroEngine): class KsayHandler(BaseHTTPRequestHandler): server_version = "ksay-kokoro/0.1" def log_message(self, fmt: str, *args: Any) -> None: sys.stderr.write("%s - %s\n" % (self.log_date_time_string(), fmt % args)) def write_json(self, status: int, value: dict[str, Any]) -> None: body = json.dumps(value).encode("utf-8") self.send_response(status) self.send_header("content-type", "application/json") self.send_header("content-length", str(len(body))) self.end_headers() self.wfile.write(body) def do_GET(self) -> None: if self.path == "/health": self.write_json( 200, { "ok": True, "model": engine.model_name, "loaded": engine.model is not None, "loadedAt": engine.loaded_at, "loadSeconds": engine.load_seconds, "defaultVoice": DEFAULT_VOICE, "defaultLangCode": DEFAULT_LANG_CODE, }, ) return self.write_json(404, {"ok": False, "error": "not found"}) def do_POST(self) -> None: if self.path != "/say": self.write_json(404, {"ok": False, "error": "not found"}) return try: payload = parse_json_body(self) result = engine.synthesize( text=str(payload.get("text", "")), voice=str(payload.get("voice") or DEFAULT_VOICE), speed=float(payload.get("speed") or 1.0), lang_code=str(payload.get("langCode") or DEFAULT_LANG_CODE), output=payload.get("output"), ) self.write_json(200, result) except Exception as exc: traceback.print_exc() self.write_json(500, {"ok": False, "error": str(exc)}) return KsayHandler def main() -> int: parser = argparse.ArgumentParser(description="Warm Kokoro TTS server for ksay.") parser.add_argument("--host", default=DEFAULT_HOST) parser.add_argument("--port", type=int, default=DEFAULT_PORT) parser.add_argument("--model", default=DEFAULT_MODEL) args = parser.parse_args() engine = KokoroEngine(args.model) print(f"Loading {args.model}...", flush=True) engine.load() print( f"ksay Kokoro ready on http://{args.host}:{args.port} " f"after {engine.load_seconds:.2f}s", flush=True, ) httpd = ThreadingHTTPServer((args.host, args.port), make_handler(engine)) def shutdown(_signum: int, _frame: Any) -> None: httpd.shutdown() signal.signal(signal.SIGTERM, shutdown) signal.signal(signal.SIGINT, shutdown) httpd.serve_forever() return 0 if __name__ == "__main__": raise SystemExit(main())