Initial TTS MCP server with 3 engines
FastMCP 3.0 Streamable HTTP server exposing Piper (Wyoming/Docker), Kokoro (ONNX), and Orpheus (Ollama+SNAC) as MCP tools. Includes a FIFO speech queue so concurrent agents don't talk over each other — waiting callers get queue position updates via ctx.info(). Tools: speak, generate_audio, list_voices, list_engines Resource: audio://recent
This commit is contained in:
commit
2fd84f0df7
17
.gitignore
vendored
Normal file
17
.gitignore
vendored
Normal file
@ -0,0 +1,17 @@
|
||||
# Models (large binary files)
|
||||
models/
|
||||
*.onnx
|
||||
*.bin
|
||||
|
||||
# Python
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.egg-info/
|
||||
dist/
|
||||
build/
|
||||
|
||||
# Virtual environment
|
||||
.venv/
|
||||
|
||||
# Environment
|
||||
.env
|
||||
195
orpheus_tts.py
Normal file
195
orpheus_tts.py
Normal file
@ -0,0 +1,195 @@
|
||||
"""Orpheus TTS via Ollama + SNAC decoder.
|
||||
|
||||
Sends text to Ollama's Orpheus model via OpenAI-compatible completions API,
|
||||
parses the <custom_token_N> responses, and decodes through SNAC to 24kHz WAV.
|
||||
|
||||
Voices: tara, leah, jess, leo, dan, mia, zac, zoe
|
||||
Emotion tags: <laugh>, <chuckle>, <sigh>, <cough>, <gasp>, <yawn>, <groan>
|
||||
"""
|
||||
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
import wave
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
|
||||
SAMPLE_RATE = 24000
|
||||
SNAC_CODEBOOK_SIZE = 4096
|
||||
|
||||
OLLAMA_URL = "http://127.0.0.1:11434"
|
||||
MODEL = "legraphista/Orpheus:3b-ft-q4_k_m"
|
||||
|
||||
TOKEN_PATTERN = re.compile(r'<custom_token_(\d+)>')
|
||||
|
||||
# Cache the SNAC model across calls
|
||||
_snac_model = None
|
||||
|
||||
|
||||
def _get_snac():
|
||||
"""Load and cache the SNAC decoder model."""
|
||||
global _snac_model
|
||||
if _snac_model is None:
|
||||
from snac import SNAC
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
print(f" Loading SNAC decoder on {device}...", file=sys.stderr)
|
||||
_snac_model = SNAC.from_pretrained("hubertsiuzdak/snac_24khz").to(device)
|
||||
return _snac_model
|
||||
|
||||
|
||||
def turn_token_into_id(token_str: str, index: int) -> int:
|
||||
"""Convert a <custom_token_N> string to a SNAC codebook ID.
|
||||
|
||||
Applies position-dependent offset: each of the 7 tokens per frame
|
||||
maps to a different SNAC codebook layer at a different offset.
|
||||
"""
|
||||
match = TOKEN_PATTERN.search(token_str)
|
||||
if not match:
|
||||
return -1
|
||||
raw_id = int(match.group(1))
|
||||
return raw_id - 10 - ((index % 7) * SNAC_CODEBOOK_SIZE)
|
||||
|
||||
|
||||
def tokens_to_audio(token_strings: list[str]) -> np.ndarray | None:
|
||||
"""Convert Orpheus custom token strings to audio via SNAC.
|
||||
|
||||
Redistributes the flat token stream into SNAC's 3 codebook layers:
|
||||
- codes_0: 1 per frame (positions 0) -> coarse
|
||||
- codes_1: 2 per frame (positions 1, 4) -> mid
|
||||
- codes_2: 4 per frame (positions 2,3,5,6) -> fine
|
||||
"""
|
||||
num_frames = len(token_strings) // 7
|
||||
if num_frames == 0:
|
||||
return None
|
||||
|
||||
token_strings = token_strings[:num_frames * 7]
|
||||
|
||||
# Convert all tokens to IDs with position-dependent offsets
|
||||
ids = [turn_token_into_id(t, i) for i, t in enumerate(token_strings)]
|
||||
|
||||
# Check for invalid tokens
|
||||
invalid = sum(1 for x in ids if x < 0 or x >= SNAC_CODEBOOK_SIZE)
|
||||
if invalid > 0:
|
||||
print(f" Warning: {invalid}/{len(ids)} invalid token IDs", file=sys.stderr)
|
||||
|
||||
# Redistribute into 3 SNAC layers
|
||||
codes_0, codes_1, codes_2 = [], [], []
|
||||
for i in range(num_frames):
|
||||
b = i * 7
|
||||
codes_0.append(ids[b + 0])
|
||||
codes_1.append(ids[b + 1])
|
||||
codes_2.append(ids[b + 2])
|
||||
codes_2.append(ids[b + 3])
|
||||
codes_1.append(ids[b + 4])
|
||||
codes_2.append(ids[b + 5])
|
||||
codes_2.append(ids[b + 6])
|
||||
|
||||
# Clamp values to valid range
|
||||
def clamp(lst):
|
||||
return [max(0, min(SNAC_CODEBOOK_SIZE - 1, x)) for x in lst]
|
||||
|
||||
codes_0 = clamp(codes_0)
|
||||
codes_1 = clamp(codes_1)
|
||||
codes_2 = clamp(codes_2)
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
snac = _get_snac()
|
||||
|
||||
codes = [
|
||||
torch.tensor(codes_0, dtype=torch.long).unsqueeze(0).to(device),
|
||||
torch.tensor(codes_1, dtype=torch.long).unsqueeze(0).to(device),
|
||||
torch.tensor(codes_2, dtype=torch.long).unsqueeze(0).to(device),
|
||||
]
|
||||
|
||||
with torch.no_grad():
|
||||
audio = snac.decode(codes)
|
||||
|
||||
return audio.squeeze().cpu().numpy()
|
||||
|
||||
|
||||
def generate_speech(text: str, voice: str = "tara") -> np.ndarray | None:
|
||||
"""Generate speech from text using Orpheus via Ollama's OpenAI-compatible API."""
|
||||
prompt = f"<|audio|>{voice}: {text}<|eot_id|>"
|
||||
|
||||
print(f" Sending to Orpheus ({voice})...", file=sys.stderr)
|
||||
t0 = time.time()
|
||||
|
||||
response = requests.post(
|
||||
f"{OLLAMA_URL}/v1/completions",
|
||||
json={
|
||||
"model": MODEL,
|
||||
"prompt": prompt,
|
||||
"max_tokens": 8192,
|
||||
"temperature": 0.6,
|
||||
"top_p": 0.9,
|
||||
"stream": False,
|
||||
},
|
||||
)
|
||||
|
||||
result = response.json()
|
||||
resp_text = result.get("choices", [{}])[0].get("text", "")
|
||||
gen_time = time.time() - t0
|
||||
|
||||
# Extract all <custom_token_N> strings
|
||||
token_strings = TOKEN_PATTERN.findall(resp_text)
|
||||
token_strings = [f"<custom_token_{t}>" for t in token_strings]
|
||||
|
||||
# Skip leading special tokens (start/header tokens with small values)
|
||||
# Audio tokens have values in range ~10-28681; tokens <10 are special
|
||||
skip = 0
|
||||
for ts in token_strings:
|
||||
m = TOKEN_PATTERN.search(ts)
|
||||
if m and int(m.group(1)) < 10:
|
||||
skip += 1
|
||||
else:
|
||||
break
|
||||
if skip > 0:
|
||||
print(f" Skipping {skip} leading special tokens", file=sys.stderr)
|
||||
token_strings = token_strings[skip:]
|
||||
|
||||
print(
|
||||
f" Got {len(token_strings)} audio tokens ({len(token_strings)//7} frames) "
|
||||
f"in {gen_time:.1f}s",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
if len(token_strings) < 7:
|
||||
print(f" Insufficient tokens. Response: {resp_text[:300]}", file=sys.stderr)
|
||||
return None
|
||||
|
||||
return tokens_to_audio(token_strings)
|
||||
|
||||
|
||||
def speak(text: str, voice: str = "tara"):
|
||||
"""Generate speech and play it through PipeWire."""
|
||||
audio = generate_speech(text, voice)
|
||||
if audio is None:
|
||||
print("Failed to generate audio", file=sys.stderr)
|
||||
return
|
||||
|
||||
# Normalize to 16-bit PCM
|
||||
peak = max(abs(audio.max()), abs(audio.min()), 1e-8)
|
||||
audio_int16 = (audio / peak * 32767).astype(np.int16)
|
||||
|
||||
tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False, prefix="orpheus-")
|
||||
with wave.open(tmp.name, "wb") as wf:
|
||||
wf.setnchannels(1)
|
||||
wf.setsampwidth(2)
|
||||
wf.setframerate(SAMPLE_RATE)
|
||||
wf.writeframes(audio_int16.tobytes())
|
||||
|
||||
duration = len(audio_int16) / SAMPLE_RATE
|
||||
print(f" Audio: {duration:.1f}s at {SAMPLE_RATE}Hz -> {tmp.name}", file=sys.stderr)
|
||||
subprocess.run(["pw-play", tmp.name])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
text = " ".join(sys.argv[1:]) or (
|
||||
"Hey Ryan, this is Orpheus speaking from the fix TTS project. "
|
||||
"Pretty wild that a language model can sound this natural, right?"
|
||||
)
|
||||
speak(text)
|
||||
36
pyproject.toml
Normal file
36
pyproject.toml
Normal file
@ -0,0 +1,36 @@
|
||||
[build-system]
|
||||
requires = ["hatchling"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "tts-mcp"
|
||||
version = "2026.02.20"
|
||||
description = "Multi-engine TTS server with speech queue, exposed via FastMCP 3.0 Streamable HTTP"
|
||||
requires-python = ">=3.12"
|
||||
license = "MIT"
|
||||
authors = [{name = "Ryan Malloy", email = "ryan@supported.systems"}]
|
||||
dependencies = [
|
||||
"fastmcp>=3.0.0",
|
||||
"kokoro-onnx>=0.5.0",
|
||||
"numpy",
|
||||
"onnxruntime",
|
||||
"pydantic-settings",
|
||||
"requests",
|
||||
"snac>=1.2.1",
|
||||
"soundfile",
|
||||
"torch",
|
||||
"wyoming>=1.8.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
tts-mcp = "tts_mcp.__main__:main"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["src/tts_mcp"]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 100
|
||||
target-version = "py312"
|
||||
|
||||
[tool.ruff.lint]
|
||||
select = ["E", "F", "I", "W"]
|
||||
3
src/tts_mcp/__init__.py
Normal file
3
src/tts_mcp/__init__.py
Normal file
@ -0,0 +1,3 @@
|
||||
"""Multi-engine TTS server with speech queue, exposed via FastMCP."""
|
||||
|
||||
__version__ = "2026.02.20"
|
||||
12
src/tts_mcp/__main__.py
Normal file
12
src/tts_mcp/__main__.py
Normal file
@ -0,0 +1,12 @@
|
||||
"""Entry point for tts-mcp server."""
|
||||
|
||||
from .server import mcp
|
||||
from .settings import settings
|
||||
|
||||
|
||||
def main():
|
||||
mcp.run(transport="streamable-http", host=settings.host, port=settings.port)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
86
src/tts_mcp/audio.py
Normal file
86
src/tts_mcp/audio.py
Normal file
@ -0,0 +1,86 @@
|
||||
"""WAV writing and audio playback utilities."""
|
||||
|
||||
import asyncio
|
||||
import itertools
|
||||
import time
|
||||
import wave
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .settings import settings
|
||||
|
||||
_counter = itertools.count(1)
|
||||
|
||||
|
||||
def write_wav(
|
||||
samples: np.ndarray,
|
||||
sample_rate: int,
|
||||
path: Path | None = None,
|
||||
prefix: str = "tts-",
|
||||
) -> Path:
|
||||
"""Write float32 or int16 samples to a WAV file.
|
||||
|
||||
If path is None, generates a timestamped filename in the output directory.
|
||||
Returns the path to the written file.
|
||||
"""
|
||||
if path is None:
|
||||
out_dir = settings.output_dir
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
ts = f"{time.strftime('%Y%m%d-%H%M%S')}-{next(_counter):04d}"
|
||||
path = out_dir / f"{prefix}{ts}.wav"
|
||||
|
||||
# Normalize float samples to 16-bit PCM
|
||||
if samples.dtype in (np.float32, np.float64):
|
||||
peak = max(abs(samples.max()), abs(samples.min()), 1e-8)
|
||||
samples = (samples / peak * 32767).astype(np.int16)
|
||||
|
||||
with wave.open(str(path), "wb") as wf:
|
||||
wf.setnchannels(1)
|
||||
wf.setsampwidth(2)
|
||||
wf.setframerate(sample_rate)
|
||||
wf.writeframes(samples.tobytes())
|
||||
|
||||
return path
|
||||
|
||||
|
||||
def write_wav_from_pcm(
|
||||
pcm_bytes: bytes,
|
||||
sample_rate: int,
|
||||
sample_width: int,
|
||||
channels: int,
|
||||
path: Path | None = None,
|
||||
prefix: str = "tts-",
|
||||
) -> Path:
|
||||
"""Write raw PCM bytes to a WAV file."""
|
||||
if path is None:
|
||||
out_dir = settings.output_dir
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
ts = f"{time.strftime('%Y%m%d-%H%M%S')}-{next(_counter):04d}"
|
||||
path = out_dir / f"{prefix}{ts}.wav"
|
||||
|
||||
with wave.open(str(path), "wb") as wf:
|
||||
wf.setnchannels(channels)
|
||||
wf.setsampwidth(sample_width)
|
||||
wf.setframerate(sample_rate)
|
||||
wf.writeframes(pcm_bytes)
|
||||
|
||||
return path
|
||||
|
||||
|
||||
def wav_duration(path: Path) -> float:
|
||||
"""Get duration of a WAV file in seconds."""
|
||||
with wave.open(str(path), "rb") as wf:
|
||||
return wf.getnframes() / wf.getframerate()
|
||||
|
||||
|
||||
async def play_audio(path: Path) -> None:
|
||||
"""Play a WAV file through PipeWire (pw-play). Async wrapper."""
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
"pw-play", str(path),
|
||||
stdout=asyncio.subprocess.DEVNULL,
|
||||
stderr=asyncio.subprocess.PIPE,
|
||||
)
|
||||
_, stderr = await proc.communicate()
|
||||
if proc.returncode != 0:
|
||||
raise RuntimeError(f"pw-play failed ({proc.returncode}): {stderr.decode().strip()}")
|
||||
1
src/tts_mcp/engines/__init__.py
Normal file
1
src/tts_mcp/engines/__init__.py
Normal file
@ -0,0 +1 @@
|
||||
"""TTS engine implementations."""
|
||||
35
src/tts_mcp/engines/base.py
Normal file
35
src/tts_mcp/engines/base.py
Normal file
@ -0,0 +1,35 @@
|
||||
"""Base classes for TTS engines."""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
@dataclass
|
||||
class TTSResult:
|
||||
"""Result of a synthesis operation."""
|
||||
|
||||
audio_path: Path
|
||||
sample_rate: int
|
||||
duration_seconds: float
|
||||
engine: str
|
||||
voice: str
|
||||
|
||||
|
||||
class TTSEngine(ABC):
|
||||
"""Abstract base for TTS engines."""
|
||||
|
||||
name: str = "unknown"
|
||||
default_voice: str = "default"
|
||||
|
||||
@abstractmethod
|
||||
async def synthesize(self, text: str, voice: str | None = None) -> TTSResult:
|
||||
"""Synthesize text to audio. Returns a TTSResult with the WAV file path."""
|
||||
|
||||
@abstractmethod
|
||||
async def list_voices(self) -> list[str]:
|
||||
"""Return available voice names."""
|
||||
|
||||
@abstractmethod
|
||||
async def check_health(self) -> dict:
|
||||
"""Check if the engine is operational. Returns status dict."""
|
||||
59
src/tts_mcp/engines/kokoro.py
Normal file
59
src/tts_mcp/engines/kokoro.py
Normal file
@ -0,0 +1,59 @@
|
||||
"""Kokoro TTS via ONNX runtime."""
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
|
||||
from kokoro_onnx import Kokoro
|
||||
|
||||
from ..audio import wav_duration, write_wav
|
||||
from ..settings import settings
|
||||
from .base import TTSEngine, TTSResult
|
||||
|
||||
|
||||
class KokoroEngine(TTSEngine):
|
||||
"""Kokoro ONNX TTS — fast local inference (~4x realtime on CPU)."""
|
||||
|
||||
name = "kokoro"
|
||||
default_voice = "af_heart"
|
||||
|
||||
def __init__(self, model: Kokoro) -> None:
|
||||
self._model = model
|
||||
|
||||
async def synthesize(self, text: str, voice: str | None = None) -> TTSResult:
|
||||
voice = voice or self.default_voice
|
||||
|
||||
# kokoro.create() is CPU-bound, run in executor
|
||||
loop = asyncio.get_running_loop()
|
||||
samples, sample_rate = await loop.run_in_executor(
|
||||
None, self._model.create, text, voice, 1.0
|
||||
)
|
||||
|
||||
path = write_wav(samples, sample_rate, prefix="kokoro-")
|
||||
|
||||
return TTSResult(
|
||||
audio_path=path,
|
||||
sample_rate=sample_rate,
|
||||
duration_seconds=wav_duration(path),
|
||||
engine=self.name,
|
||||
voice=voice,
|
||||
)
|
||||
|
||||
async def list_voices(self) -> list[str]:
|
||||
try:
|
||||
voices = self._model.get_voices()
|
||||
blacklist = settings.blacklisted_voices
|
||||
return sorted(v for v in voices if v.lower() not in blacklist)
|
||||
except Exception as e:
|
||||
print(f"Kokoro voice listing failed: {e}", file=sys.stderr)
|
||||
return []
|
||||
|
||||
async def check_health(self) -> dict:
|
||||
try:
|
||||
voices = self._model.get_voices()
|
||||
return {
|
||||
"status": "healthy",
|
||||
"engine": self.name,
|
||||
"voice_count": len(voices),
|
||||
}
|
||||
except Exception as e:
|
||||
return {"status": "unhealthy", "engine": self.name, "error": str(e)}
|
||||
190
src/tts_mcp/engines/orpheus.py
Normal file
190
src/tts_mcp/engines/orpheus.py
Normal file
@ -0,0 +1,190 @@
|
||||
"""Orpheus TTS via Ollama + SNAC decoder.
|
||||
|
||||
Sends text to Ollama's Orpheus model, parses <custom_token_N> responses,
|
||||
and decodes through SNAC to 24kHz WAV. SNAC runs on CPU (RTX 5070 SM 120
|
||||
is not yet supported by PyTorch 2.6).
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import requests
|
||||
import torch
|
||||
|
||||
from ..audio import wav_duration, write_wav
|
||||
from ..settings import settings
|
||||
from .base import TTSEngine, TTSResult
|
||||
|
||||
SAMPLE_RATE = 24000
|
||||
SNAC_CODEBOOK_SIZE = 4096
|
||||
TOKEN_PATTERN = re.compile(r"<custom_token_(\d+)>")
|
||||
|
||||
ALL_VOICES = ["tara", "leah", "jess", "leo", "dan", "mia", "zac", "zoe"]
|
||||
|
||||
|
||||
def _turn_token_into_id(token_str: str, index: int) -> int:
|
||||
"""Convert a <custom_token_N> string to a SNAC codebook ID.
|
||||
|
||||
Applies position-dependent offset: each of the 7 tokens per frame
|
||||
maps to a different SNAC codebook layer at a different offset.
|
||||
"""
|
||||
match = TOKEN_PATTERN.search(token_str)
|
||||
if not match:
|
||||
return -1
|
||||
raw_id = int(match.group(1))
|
||||
return raw_id - 10 - ((index % 7) * SNAC_CODEBOOK_SIZE)
|
||||
|
||||
|
||||
def _tokens_to_audio(token_strings: list[str], snac_model) -> np.ndarray | None:
|
||||
"""Convert Orpheus custom token strings to audio via SNAC.
|
||||
|
||||
Redistributes the flat token stream into SNAC's 3 codebook layers:
|
||||
- codes_0: 1 per frame (positions 0) -> coarse
|
||||
- codes_1: 2 per frame (positions 1, 4) -> mid
|
||||
- codes_2: 4 per frame (positions 2,3,5,6) -> fine
|
||||
"""
|
||||
num_frames = len(token_strings) // 7
|
||||
if num_frames == 0:
|
||||
return None
|
||||
|
||||
token_strings = token_strings[: num_frames * 7]
|
||||
ids = [_turn_token_into_id(t, i) for i, t in enumerate(token_strings)]
|
||||
|
||||
invalid = sum(1 for x in ids if x < 0 or x >= SNAC_CODEBOOK_SIZE)
|
||||
if invalid > 0:
|
||||
print(f" Warning: {invalid}/{len(ids)} invalid token IDs", file=sys.stderr)
|
||||
|
||||
codes_0, codes_1, codes_2 = [], [], []
|
||||
for i in range(num_frames):
|
||||
b = i * 7
|
||||
codes_0.append(ids[b + 0])
|
||||
codes_1.append(ids[b + 1])
|
||||
codes_2.append(ids[b + 2])
|
||||
codes_2.append(ids[b + 3])
|
||||
codes_1.append(ids[b + 4])
|
||||
codes_2.append(ids[b + 5])
|
||||
codes_2.append(ids[b + 6])
|
||||
|
||||
def clamp(lst):
|
||||
return [max(0, min(SNAC_CODEBOOK_SIZE - 1, x)) for x in lst]
|
||||
|
||||
# SNAC runs on CPU — RTX 5070 SM 120 not supported by PyTorch 2.6
|
||||
device = "cpu"
|
||||
codes = [
|
||||
torch.tensor(clamp(codes_0), dtype=torch.long).unsqueeze(0).to(device),
|
||||
torch.tensor(clamp(codes_1), dtype=torch.long).unsqueeze(0).to(device),
|
||||
torch.tensor(clamp(codes_2), dtype=torch.long).unsqueeze(0).to(device),
|
||||
]
|
||||
|
||||
with torch.no_grad():
|
||||
audio = snac_model.decode(codes)
|
||||
|
||||
return audio.squeeze().cpu().numpy()
|
||||
|
||||
|
||||
class OrpheusEngine(TTSEngine):
|
||||
"""Orpheus TTS via Ollama's completions API + SNAC audio decoding."""
|
||||
|
||||
name = "orpheus"
|
||||
default_voice = "tara"
|
||||
|
||||
def __init__(self, snac_model, ollama_url: str, model_name: str) -> None:
|
||||
self._snac = snac_model
|
||||
self._ollama_url = ollama_url
|
||||
self._model = model_name
|
||||
|
||||
async def synthesize(self, text: str, voice: str | None = None) -> TTSResult:
|
||||
voice = voice or self.default_voice
|
||||
prompt = f"<|audio|>{voice}: {text}<|eot_id|>"
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
# Ollama API call (blocking HTTP)
|
||||
t0 = time.time()
|
||||
|
||||
def _call_ollama():
|
||||
resp = requests.post(
|
||||
f"{self._ollama_url}/v1/completions",
|
||||
json={
|
||||
"model": self._model,
|
||||
"prompt": prompt,
|
||||
"max_tokens": 8192,
|
||||
"temperature": 0.6,
|
||||
"top_p": 0.9,
|
||||
"stream": False,
|
||||
},
|
||||
timeout=120,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
result = await loop.run_in_executor(None, _call_ollama)
|
||||
gen_time = time.time() - t0
|
||||
|
||||
resp_text = result.get("choices", [{}])[0].get("text", "")
|
||||
|
||||
# Extract <custom_token_N> strings
|
||||
token_strings = TOKEN_PATTERN.findall(resp_text)
|
||||
token_strings = [f"<custom_token_{t}>" for t in token_strings]
|
||||
|
||||
# Skip leading special tokens (values < 10)
|
||||
skip = 0
|
||||
for ts in token_strings:
|
||||
m = TOKEN_PATTERN.search(ts)
|
||||
if m and int(m.group(1)) < 10:
|
||||
skip += 1
|
||||
else:
|
||||
break
|
||||
if skip > 0:
|
||||
token_strings = token_strings[skip:]
|
||||
|
||||
print(
|
||||
f" Orpheus: {len(token_strings)} tokens "
|
||||
f"({len(token_strings) // 7} frames) in {gen_time:.1f}s",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
if len(token_strings) < 7:
|
||||
raise RuntimeError(
|
||||
f"Orpheus returned insufficient tokens ({len(token_strings)}). "
|
||||
f"Response preview: {resp_text[:200]}"
|
||||
)
|
||||
|
||||
# SNAC decode (CPU-bound)
|
||||
audio = await loop.run_in_executor(
|
||||
None, _tokens_to_audio, token_strings, self._snac
|
||||
)
|
||||
if audio is None:
|
||||
raise RuntimeError("SNAC decoding produced no audio")
|
||||
|
||||
path = write_wav(audio, SAMPLE_RATE, prefix="orpheus-")
|
||||
|
||||
return TTSResult(
|
||||
audio_path=path,
|
||||
sample_rate=SAMPLE_RATE,
|
||||
duration_seconds=wav_duration(path),
|
||||
engine=self.name,
|
||||
voice=voice,
|
||||
)
|
||||
|
||||
async def list_voices(self) -> list[str]:
|
||||
blacklist = settings.blacklisted_voices
|
||||
return sorted(v for v in ALL_VOICES if v.lower() not in blacklist)
|
||||
|
||||
async def check_health(self) -> dict:
|
||||
try:
|
||||
resp = requests.get(f"{self._ollama_url}/api/tags", timeout=5)
|
||||
resp.raise_for_status()
|
||||
models = [m["name"] for m in resp.json().get("models", [])]
|
||||
has_orpheus = any("orpheus" in m.lower() for m in models)
|
||||
return {
|
||||
"status": "healthy" if has_orpheus else "degraded",
|
||||
"engine": self.name,
|
||||
"model_loaded": has_orpheus,
|
||||
"ollama_models": len(models),
|
||||
}
|
||||
except Exception as e:
|
||||
return {"status": "unhealthy", "engine": self.name, "error": str(e)}
|
||||
101
src/tts_mcp/engines/piper.py
Normal file
101
src/tts_mcp/engines/piper.py
Normal file
@ -0,0 +1,101 @@
|
||||
"""Piper TTS via the Wyoming protocol over TCP."""
|
||||
|
||||
import sys
|
||||
|
||||
from wyoming.audio import AudioChunk, AudioStart, AudioStop
|
||||
from wyoming.client import AsyncTcpClient
|
||||
from wyoming.info import Describe, Info
|
||||
from wyoming.tts import Synthesize, SynthesizeVoice
|
||||
|
||||
from ..audio import wav_duration, write_wav_from_pcm
|
||||
from ..settings import settings
|
||||
from .base import TTSEngine, TTSResult
|
||||
|
||||
|
||||
class PiperEngine(TTSEngine):
|
||||
"""Piper TTS accessed via Wyoming protocol (Docker container)."""
|
||||
|
||||
name = "piper"
|
||||
default_voice = "en_US-lessac-medium"
|
||||
|
||||
def __init__(self, host: str, port: int) -> None:
|
||||
self._host = host
|
||||
self._port = port
|
||||
|
||||
async def synthesize(self, text: str, voice: str | None = None) -> TTSResult:
|
||||
voice = voice or self.default_voice
|
||||
|
||||
async with AsyncTcpClient(self._host, self._port) as client:
|
||||
# Send synthesize request
|
||||
synth = Synthesize(text=text, voice=SynthesizeVoice(name=voice))
|
||||
await client.write_event(synth.event())
|
||||
|
||||
# Collect audio chunks
|
||||
pcm_chunks: list[bytes] = []
|
||||
sample_rate = 22050
|
||||
sample_width = 2
|
||||
channels = 1
|
||||
|
||||
while True:
|
||||
event = await client.read_event()
|
||||
if event is None:
|
||||
break
|
||||
|
||||
if AudioStart.is_type(event.type):
|
||||
start = AudioStart.from_event(event)
|
||||
sample_rate = start.rate
|
||||
sample_width = start.width
|
||||
channels = start.channels
|
||||
|
||||
elif AudioChunk.is_type(event.type):
|
||||
chunk = AudioChunk.from_event(event)
|
||||
pcm_chunks.append(chunk.audio)
|
||||
|
||||
elif AudioStop.is_type(event.type):
|
||||
break
|
||||
|
||||
pcm_data = b"".join(pcm_chunks)
|
||||
if not pcm_data:
|
||||
raise RuntimeError("Piper returned no audio data")
|
||||
|
||||
path = write_wav_from_pcm(
|
||||
pcm_data, sample_rate, sample_width, channels, prefix="piper-"
|
||||
)
|
||||
|
||||
return TTSResult(
|
||||
audio_path=path,
|
||||
sample_rate=sample_rate,
|
||||
duration_seconds=wav_duration(path),
|
||||
engine=self.name,
|
||||
voice=voice,
|
||||
)
|
||||
|
||||
async def list_voices(self) -> list[str]:
|
||||
try:
|
||||
async with AsyncTcpClient(self._host, self._port) as client:
|
||||
await client.write_event(Describe().event())
|
||||
event = await client.read_event()
|
||||
if event is None or not Info.is_type(event.type):
|
||||
return []
|
||||
|
||||
info = Info.from_event(event)
|
||||
voices = []
|
||||
for tts_prog in info.tts:
|
||||
for v in tts_prog.voices:
|
||||
if v.name.lower() not in settings.blacklisted_voices:
|
||||
voices.append(v.name)
|
||||
return sorted(voices)
|
||||
except (OSError, ConnectionError) as e:
|
||||
print(f"Piper voice listing failed: {e}", file=sys.stderr)
|
||||
return []
|
||||
|
||||
async def check_health(self) -> dict:
|
||||
try:
|
||||
async with AsyncTcpClient(self._host, self._port) as client:
|
||||
await client.write_event(Describe().event())
|
||||
event = await client.read_event()
|
||||
if event and Info.is_type(event.type):
|
||||
return {"status": "healthy", "engine": self.name}
|
||||
except (OSError, ConnectionError) as e:
|
||||
return {"status": "unhealthy", "engine": self.name, "error": str(e)}
|
||||
return {"status": "unhealthy", "engine": self.name, "error": "no response"}
|
||||
110
src/tts_mcp/queue.py
Normal file
110
src/tts_mcp/queue.py
Normal file
@ -0,0 +1,110 @@
|
||||
"""Speech queue — serializes playback so agents don't talk over each other."""
|
||||
|
||||
import asyncio
|
||||
from collections import deque
|
||||
from contextlib import asynccontextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import AsyncIterator
|
||||
|
||||
from .audio import play_audio
|
||||
from .engines.base import TTSResult
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Waiter:
|
||||
caller_id: str
|
||||
text_preview: str
|
||||
|
||||
|
||||
class SpeechQueue:
|
||||
"""FIFO queue for audio playback.
|
||||
|
||||
Only one audio file plays at a time. Callers that arrive while
|
||||
someone is speaking wait in line and get progress updates via
|
||||
their MCP context.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._lock = asyncio.Lock()
|
||||
self._current: _Waiter | None = None
|
||||
self._waiters: deque[_Waiter] = deque()
|
||||
self._counter = 0
|
||||
|
||||
def _next_id(self) -> str:
|
||||
self._counter += 1
|
||||
return f"speaker-{self._counter}"
|
||||
|
||||
@property
|
||||
def depth(self) -> int:
|
||||
return len(self._waiters)
|
||||
|
||||
@property
|
||||
def current_speaker(self) -> str | None:
|
||||
return self._current.caller_id if self._current else None
|
||||
|
||||
def status(self) -> dict:
|
||||
return {
|
||||
"current_speaker": self.current_speaker,
|
||||
"queue_depth": self.depth,
|
||||
"waiting": [w.caller_id for w in self._waiters],
|
||||
}
|
||||
|
||||
@asynccontextmanager
|
||||
async def acquire(
|
||||
self,
|
||||
caller_id: str | None = None,
|
||||
text_preview: str = "",
|
||||
) -> AsyncIterator[str]:
|
||||
"""Context manager that waits for the speaker's turn.
|
||||
|
||||
Yields the caller_id once it's this caller's turn to play audio.
|
||||
"""
|
||||
cid = caller_id or self._next_id()
|
||||
waiter = _Waiter(caller_id=cid, text_preview=text_preview[:60])
|
||||
self._waiters.append(waiter)
|
||||
|
||||
try:
|
||||
async with self._lock:
|
||||
# We're up — remove ourselves from the waiting list
|
||||
if waiter in self._waiters:
|
||||
self._waiters.remove(waiter)
|
||||
self._current = waiter
|
||||
yield cid
|
||||
finally:
|
||||
if self._current is waiter:
|
||||
self._current = None
|
||||
|
||||
async def speak(
|
||||
self,
|
||||
result: TTSResult,
|
||||
caller_id: str | None = None,
|
||||
info_callback=None,
|
||||
) -> dict:
|
||||
"""Queue and play a TTSResult. Returns status dict when done.
|
||||
|
||||
info_callback: async callable(message) for progress updates (e.g. ctx.info).
|
||||
"""
|
||||
cid = caller_id or self._next_id()
|
||||
preview = f"{result.engine}/{result.voice}"
|
||||
|
||||
# Show queue position before acquiring
|
||||
if self._lock.locked():
|
||||
pos = self.depth + 1
|
||||
msg = f"Queued at position {pos}"
|
||||
if self._current:
|
||||
msg += f" (currently playing: {self._current.caller_id})"
|
||||
if info_callback:
|
||||
await info_callback(msg)
|
||||
|
||||
async with self.acquire(cid, preview):
|
||||
if info_callback:
|
||||
await info_callback(f"Now playing: {result.engine}/{result.voice}")
|
||||
await play_audio(result.audio_path)
|
||||
|
||||
return {
|
||||
"played": True,
|
||||
"file": str(result.audio_path),
|
||||
"duration_seconds": result.duration_seconds,
|
||||
"engine": result.engine,
|
||||
"voice": result.voice,
|
||||
}
|
||||
243
src/tts_mcp/server.py
Normal file
243
src/tts_mcp/server.py
Normal file
@ -0,0 +1,243 @@
|
||||
"""FastMCP 3.0 server — tools, lifespan, and resource definitions."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Literal
|
||||
|
||||
from fastmcp import Context, FastMCP
|
||||
from fastmcp.server.dependencies import CurrentContext
|
||||
|
||||
from .engines.base import TTSEngine
|
||||
from .engines.kokoro import KokoroEngine
|
||||
from .engines.orpheus import OrpheusEngine
|
||||
from .engines.piper import PiperEngine
|
||||
from .queue import SpeechQueue
|
||||
from .settings import settings
|
||||
|
||||
ENGINE_NAMES = Literal["piper", "kokoro", "orpheus"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Lifespan — load models once at startup, share across all requests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@asynccontextmanager
|
||||
async def app_lifespan(server: FastMCP):
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
# --- Kokoro (ONNX) ---
|
||||
print("Loading Kokoro ONNX model...", file=sys.stderr)
|
||||
t0 = time.time()
|
||||
from kokoro_onnx import Kokoro
|
||||
|
||||
kokoro_model = await loop.run_in_executor(
|
||||
None,
|
||||
Kokoro,
|
||||
str(settings.kokoro_model),
|
||||
str(settings.kokoro_voices),
|
||||
)
|
||||
print(f" Kokoro ready in {time.time() - t0:.1f}s", file=sys.stderr)
|
||||
|
||||
# --- SNAC decoder (CPU only — RTX 5070 SM 120 unsupported by PyTorch 2.6) ---
|
||||
print("Loading SNAC decoder on CPU...", file=sys.stderr)
|
||||
t0 = time.time()
|
||||
|
||||
def _load_snac():
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||||
from snac import SNAC
|
||||
|
||||
return SNAC.from_pretrained("hubertsiuzdak/snac_24khz").to("cpu")
|
||||
|
||||
snac_model = await loop.run_in_executor(None, _load_snac)
|
||||
print(f" SNAC ready in {time.time() - t0:.1f}s", file=sys.stderr)
|
||||
|
||||
# --- Build engines ---
|
||||
engines: dict[str, TTSEngine] = {
|
||||
"piper": PiperEngine(settings.piper_host, settings.piper_port),
|
||||
"kokoro": KokoroEngine(kokoro_model),
|
||||
"orpheus": OrpheusEngine(snac_model, settings.ollama_url, settings.orpheus_model),
|
||||
}
|
||||
|
||||
# Health check all engines at startup
|
||||
for name, eng in engines.items():
|
||||
health = await eng.check_health()
|
||||
print(f" {name}: {health['status']}", file=sys.stderr)
|
||||
|
||||
queue = SpeechQueue()
|
||||
|
||||
print(
|
||||
f"TTS MCP server ready on {settings.host}:{settings.port} "
|
||||
f"with {len(engines)} engines",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
try:
|
||||
yield {"engines": engines, "queue": queue}
|
||||
finally:
|
||||
print("TTS MCP server shutting down", file=sys.stderr)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FastMCP instance
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
mcp = FastMCP(
|
||||
"tts-mcp",
|
||||
instructions=(
|
||||
"Multi-engine text-to-speech server. Use 'speak' to synthesize and play audio "
|
||||
"through the host speakers (queued so agents don't talk over each other). "
|
||||
"Use 'generate_audio' to synthesize without playing. "
|
||||
"Engines: kokoro (fast ONNX, ~50 voices), piper (Wyoming/Docker), "
|
||||
"orpheus (Ollama LLM, supports <laugh> etc.)."
|
||||
),
|
||||
lifespan=app_lifespan,
|
||||
)
|
||||
|
||||
|
||||
def _get_state(ctx: Context) -> tuple[dict[str, TTSEngine], SpeechQueue]:
|
||||
"""Extract engines and queue from lifespan context."""
|
||||
state = ctx.lifespan_context
|
||||
return state["engines"], state["queue"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@mcp.tool
|
||||
async def speak(
|
||||
text: str,
|
||||
engine: ENGINE_NAMES = "kokoro",
|
||||
voice: str | None = None,
|
||||
ctx: Context = CurrentContext(),
|
||||
) -> dict:
|
||||
"""Synthesize text and play it through the host speakers.
|
||||
|
||||
Audio is queued — if another agent is currently speaking, you'll wait
|
||||
your turn and get notified when playback starts.
|
||||
|
||||
Args:
|
||||
text: Text to speak. Orpheus supports emotion tags like <laugh>, <sigh>, etc.
|
||||
engine: TTS engine to use. kokoro is fastest, orpheus is most expressive.
|
||||
voice: Voice name (use list_voices to see options). None = engine default.
|
||||
"""
|
||||
engines, queue = _get_state(ctx)
|
||||
|
||||
if engine not in engines:
|
||||
return {"error": f"Unknown engine: {engine}. Available: {list(engines.keys())}"}
|
||||
|
||||
eng = engines[engine]
|
||||
|
||||
# Synthesize audio (not queued — multiple agents can synthesize simultaneously)
|
||||
await ctx.info(f"Synthesizing with {engine}...")
|
||||
result = await eng.synthesize(text, voice)
|
||||
|
||||
# Queue for playback (serialized)
|
||||
return await queue.speak(
|
||||
result,
|
||||
info_callback=ctx.info,
|
||||
)
|
||||
|
||||
|
||||
@mcp.tool
|
||||
async def generate_audio(
|
||||
text: str,
|
||||
engine: ENGINE_NAMES = "kokoro",
|
||||
voice: str | None = None,
|
||||
ctx: Context = CurrentContext(),
|
||||
) -> dict:
|
||||
"""Synthesize text to a WAV file without playing it.
|
||||
|
||||
Bypasses the speech queue — multiple agents can generate simultaneously.
|
||||
Returns the file path and metadata.
|
||||
|
||||
Args:
|
||||
text: Text to synthesize.
|
||||
engine: TTS engine to use.
|
||||
voice: Voice name (use list_voices to see options). None = engine default.
|
||||
"""
|
||||
engines, _ = _get_state(ctx)
|
||||
|
||||
if engine not in engines:
|
||||
return {"error": f"Unknown engine: {engine}. Available: {list(engines.keys())}"}
|
||||
|
||||
eng = engines[engine]
|
||||
await ctx.info(f"Generating audio with {engine}...")
|
||||
result = await eng.synthesize(text, voice)
|
||||
|
||||
return {
|
||||
"file": str(result.audio_path),
|
||||
"duration_seconds": result.duration_seconds,
|
||||
"sample_rate": result.sample_rate,
|
||||
"engine": result.engine,
|
||||
"voice": result.voice,
|
||||
}
|
||||
|
||||
|
||||
@mcp.tool
|
||||
async def list_voices(
|
||||
engine: ENGINE_NAMES,
|
||||
ctx: Context = CurrentContext(),
|
||||
) -> list[str]:
|
||||
"""List available voices for a TTS engine.
|
||||
|
||||
Blacklisted voices are excluded. Use the returned names as the
|
||||
'voice' parameter in speak/generate_audio.
|
||||
|
||||
Args:
|
||||
engine: Which engine to list voices for.
|
||||
"""
|
||||
engines, _ = _get_state(ctx)
|
||||
|
||||
if engine not in engines:
|
||||
return []
|
||||
|
||||
return await engines[engine].list_voices()
|
||||
|
||||
|
||||
@mcp.tool
|
||||
async def list_engines(
|
||||
ctx: Context = CurrentContext(),
|
||||
) -> list[dict]:
|
||||
"""Show all TTS engines and their health status.
|
||||
|
||||
Returns engine name, default voice, and health check results.
|
||||
"""
|
||||
engines, queue = _get_state(ctx)
|
||||
|
||||
results = []
|
||||
for name, eng in engines.items():
|
||||
health = await eng.check_health()
|
||||
results.append({
|
||||
"engine": name,
|
||||
"default_voice": eng.default_voice,
|
||||
**health,
|
||||
})
|
||||
|
||||
# Include queue status
|
||||
results.append({"queue": queue.status()})
|
||||
return results
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Resource — recent audio files
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@mcp.resource("audio://recent")
|
||||
async def recent_audio() -> str:
|
||||
"""List recently generated audio files."""
|
||||
out_dir = settings.output_dir
|
||||
if not out_dir.exists():
|
||||
return "No audio files yet."
|
||||
|
||||
wavs = sorted(out_dir.glob("*.wav"), key=lambda p: p.stat().st_mtime, reverse=True)
|
||||
lines = []
|
||||
for w in wavs[:20]:
|
||||
size_kb = w.stat().st_size / 1024
|
||||
mtime = time.strftime("%H:%M:%S", time.localtime(w.stat().st_mtime))
|
||||
lines.append(f"{mtime} {size_kb:6.1f}KB {w.name}")
|
||||
|
||||
return "\n".join(lines) if lines else "No audio files yet."
|
||||
46
src/tts_mcp/settings.py
Normal file
46
src/tts_mcp/settings.py
Normal file
@ -0,0 +1,46 @@
|
||||
"""Configuration loaded from environment / .env file."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from pydantic_settings import BaseSettings
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = {"env_prefix": "TTS_", "env_file": ".env", "extra": "ignore"}
|
||||
|
||||
# Server
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 8371
|
||||
|
||||
# Piper (Wyoming protocol)
|
||||
piper_host: str = "172.26.0.3"
|
||||
piper_port: int = 10200
|
||||
|
||||
# Kokoro (ONNX)
|
||||
kokoro_model: Path = Path("models/kokoro/kokoro-v1.0.onnx")
|
||||
kokoro_voices: Path = Path("models/kokoro/voices-v1.0.bin")
|
||||
|
||||
# Orpheus (Ollama + SNAC)
|
||||
ollama_url: str = "http://127.0.0.1:11434"
|
||||
orpheus_model: str = "legraphista/Orpheus:3b-ft-q4_k_m"
|
||||
|
||||
# Voice filtering
|
||||
voice_blacklist: str = "amy,jess,zoe,adam"
|
||||
|
||||
# Audio output (empty = system temp dir)
|
||||
audio_dir: str = ""
|
||||
|
||||
@property
|
||||
def blacklisted_voices(self) -> set[str]:
|
||||
return {v.strip().lower() for v in self.voice_blacklist.split(",") if v.strip()}
|
||||
|
||||
@property
|
||||
def output_dir(self) -> Path:
|
||||
if self.audio_dir:
|
||||
p = Path(self.audio_dir)
|
||||
p.mkdir(parents=True, exist_ok=True)
|
||||
return p
|
||||
return Path("/tmp/tts-mcp")
|
||||
|
||||
|
||||
settings = Settings()
|
||||
Loading…
x
Reference in New Issue
Block a user