#!/usr/bin/env python3
"""
generate_speech.py - minimal Arabic TTS test using Piper (via sherpa-onnx)

Setup (see README.md for full instructions):
    pip install sherpa-onnx --break-system-packages
    wget https://github.com/k2-fsa/sherpa-onnx/releases/download/tts-models/vits-piper-ar_JO-kareem-medium.tar.bz2
    tar xf vits-piper-ar_JO-kareem-medium.tar.bz2

Usage:
    python3 generate_speech.py "مرحباً، كيف يمكنني مساعدتك اليوم؟" output.wav
"""
import sys
import wave
import numpy as np
import sherpa_onnx

MODEL_DIR = "vits-piper-ar_JO-kareem-medium"


def main():
    if len(sys.argv) < 2:
        print("Usage: python3 generate_speech.py \"Arabic text\" [output.wav]")
        sys.exit(1)

    text = sys.argv[1]
    output_path = sys.argv[2] if len(sys.argv) > 2 else "output.wav"

    tts = sherpa_onnx.OfflineTts(
        sherpa_onnx.OfflineTtsConfig(
            model=sherpa_onnx.OfflineTtsModelConfig(
                vits=sherpa_onnx.OfflineTtsVitsModelConfig(
                    model=f"{MODEL_DIR}/ar_JO-kareem-medium.onnx",
                    lexicon="",
                    tokens=f"{MODEL_DIR}/tokens.txt",
                    data_dir=f"{MODEL_DIR}/espeak-ng-data",
                ),
                num_threads=2,
            ),
        )
    )

    audio = tts.generate(text, sid=0, speed=1.0)

    with wave.open(output_path, "wb") as f:
        f.setnchannels(1)
        f.setsampwidth(2)
        f.setframerate(audio.sample_rate)
        samples = (np.array(audio.samples) * 32767).astype(np.int16)
        f.writeframes(samples.tobytes())

    print(f"Generated: {output_path} (sample rate: {audio.sample_rate}Hz)")


if __name__ == "__main__":
    main()
