astra-local-voice/astra/server.py
Jeuner a42b3e8d73 feat: initial commit of astra local voice agent
Lokaler deutscher Sprachagent für Apple Silicon: Pipecat-Pipeline mit
Nemotron-ASR (MLX), Qwen über natives Ollama /api/chat, und Pocket TTS.
Loopback-only WebRTC-Server mit Origin/Host-Härtung.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01LVgSHNHdRx3UNTBodFmhRA
2026-09-07 13:39:51 +02:00

322 lines
12 KiB
Python

"""Loopback-only WebRTC voice app. One live session shares the warm models."""
import asyncio
import contextlib
import os
import sys
from collections import deque
from contextlib import asynccontextmanager
from pathlib import Path
from typing import Literal
import httpx
import uvicorn
from fastapi import FastAPI, HTTPException, Request
from fastapi.responses import FileResponse, JSONResponse
from fastapi.staticfiles import StaticFiles
from loguru import logger
from pydantic import BaseModel, Field
from astra.core import SYSTEM_PROMPT, Settings, build_request, local_origin_allowed
from astra.inference import Models, on_executor
ROOT = Path(__file__).resolve().parent.parent
class Offer(BaseModel):
sdp: str = Field(min_length=10, max_length=65536)
type: Literal["offer"]
class Disconnect(BaseModel):
pc_id: str = Field(max_length=100)
async def run_voice(connection, models, config):
from pipecat.audio.vad.silero import SileroVADAnalyzer
from pipecat.audio.vad.vad_analyzer import VADParams
from pipecat.frames.frames import (
BotStartedSpeakingFrame,
BotStoppedSpeakingFrame,
InterruptionFrame,
LLMFullResponseStartFrame,
VADUserStartedSpeakingFrame,
)
from pipecat.observers.base_observer import BaseObserver
from pipecat.pipeline.pipeline import Pipeline
from pipecat.pipeline.worker import PipelineParams, PipelineWorker
from pipecat.processors.aggregators.llm_context import LLMContext
from pipecat.processors.aggregators.llm_response_universal import (
LLMContextAggregatorPair,
LLMUserAggregatorParams,
)
from pipecat.transports.base_transport import TransportParams
from pipecat.transports.smallwebrtc.transport import SmallWebRTCTransport
from pipecat.workers.runner import WorkerRunner
from astra.services import LocalPocketTTSService, NativeOllamaService, NemotronSTTService
notify = connection.send_app_message
class UIObserver(BaseObserver):
def __init__(self):
super().__init__()
self.seen = deque(maxlen=128)
async def on_push_frame(self, data):
frame = data.frame
relevant = (
BotStartedSpeakingFrame,
BotStoppedSpeakingFrame,
InterruptionFrame,
LLMFullResponseStartFrame,
VADUserStartedSpeakingFrame,
)
if not isinstance(frame, relevant) or frame.id in self.seen:
return
self.seen.append(frame.id)
if isinstance(frame, BotStartedSpeakingFrame):
notify({"type": "state", "state": "speaking"})
elif isinstance(frame, LLMFullResponseStartFrame):
notify({"type": "state", "state": "responding"})
else:
notify({"type": "state", "state": "listening"})
transport = SmallWebRTCTransport(
connection,
TransportParams(
audio_in_enabled=True,
audio_out_enabled=True,
audio_in_sample_rate=16000,
audio_out_sample_rate=24000,
),
)
context = LLMContext([{"role": "system", "content": SYSTEM_PROMPT}])
aggregators = LLMContextAggregatorPair(
context,
user_params=LLMUserAggregatorParams(
vad_analyzer=SileroVADAnalyzer(
params=VADParams(
confidence=0.7,
start_secs=0.1,
stop_secs=0.7,
min_volume=0.6,
)
),
),
)
stt = NemotronSTTService(models, notify)
llm = NativeOllamaService(config, notify)
tts = LocalPocketTTSService(models)
pipeline = Pipeline(
[
transport.input(),
stt,
aggregators.user(),
llm,
tts,
transport.output(),
aggregators.assistant(),
]
)
worker = PipelineWorker(
pipeline,
params=PipelineParams(audio_in_sample_rate=16000, audio_out_sample_rate=24000),
enable_rtvi=False,
idle_timeout_secs=None,
observers=[UIObserver()],
)
@transport.event_handler("on_client_connected")
async def connected(transport, client):
notify({"type": "state", "state": "listening"})
@transport.event_handler("on_client_disconnected")
async def disconnected(transport, client):
await worker.cancel()
@aggregators.user().event_handler("on_user_turn_stopped")
async def user_turn(aggregator, strategy, message):
if message.content:
notify({"type": "transcript", "role": "user", "text": message.content})
@aggregators.assistant().event_handler("on_assistant_turn_stopped")
async def assistant_turn(aggregator, message):
if message.content:
notify(
{
"type": "transcript",
"role": "assistant",
"text": message.content,
"interrupted": message.interrupted,
}
)
@worker.event_handler("on_pipeline_error")
async def error(worker, frame):
notify({"type": "error", "message": frame.error})
runner = WorkerRunner(handle_sigint=False)
await runner.add_workers(worker)
await runner.run()
def create_app(config=None, *, load_models=True):
config = config or Settings.from_env()
models = Models(config)
status = {"ready": False, "stage": "starting", "error": None}
sessions = {}
lock = asyncio.Lock()
async def warmup():
try:
status["stage"] = "Gesprächssteuerung wird vorbereitet"
def prepare_pipeline():
import nltk
from pipecat.audio.turn.smart_turn.local_smart_turn_v3 import (
LocalSmartTurnAnalyzerV3,
)
from pipecat.audio.vad.silero import SileroVADAnalyzer
from astra import services # noqa: F401
nltk.data.path.insert(0, str(ROOT / ".cache" / "nltk"))
nltk.data.find("tokenizers/punkt_tab")
nltk.sent_tokenize("Hallo. Alles bereit.")
LocalSmartTurnAnalyzerV3()
SileroVADAnalyzer()
await asyncio.to_thread(prepare_pipeline)
status["stage"] = "Spracherkennung wird geladen"
await on_executor(models.stt_executor, models.load_stt)
status["stage"] = "Deutsche Stimme wird geladen"
await on_executor(models.tts_executor, models.load_tts)
status["stage"] = "Qwen wird vorbereitet"
async with httpx.AsyncClient(timeout=180) as client:
payload = build_request(config, [{"role": "user", "content": "Sage Hallo."}])
payload["stream"] = False
payload["options"]["num_predict"] = 12
response = await client.post(f"{config.ollama_url}/api/chat", json=payload)
response.raise_for_status()
result = response.json()
if result.get("error") or not result.get("message", {}).get("content"):
raise RuntimeError(f"Ollama: {result.get('error', 'Keine Antwort')}")
if result["message"].get("thinking"):
raise RuntimeError("Ollama hat Thinking nicht ausgeschaltet")
status.update(ready=True, stage="Bereit")
logger.info("Astra bereit auf http://localhost:{}", config.port)
except asyncio.CancelledError:
raise
except Exception as exc:
logger.exception("Modelle konnten nicht geladen werden")
message = str(exc) or type(exc).__name__
if isinstance(exc, httpx.TimeoutException):
message = "Ollama antwortet nicht rechtzeitig. Bitte den lokalen Modelldienst prüfen."
status.update(ready=False, stage="Start fehlgeschlagen", error=message)
@asynccontextmanager
async def lifespan(app):
warming = asyncio.create_task(warmup()) if load_models else None
try:
yield
finally:
if warming:
warming.cancel()
with contextlib.suppress(asyncio.CancelledError):
await warming
for connection, task in list(sessions.values()):
await connection.disconnect()
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
await models.close()
app = FastAPI(lifespan=lifespan, docs_url=None, redoc_url=None, openapi_url=None)
app.state.status = status
@app.middleware("http")
async def local_only(request: Request, call_next):
if request.url.hostname not in {"localhost", "127.0.0.1"}:
return JSONResponse({"detail": "Nur lokal erreichbar"}, status_code=403)
if request.method == "POST" and not local_origin_allowed(
request.headers.get("origin", ""), config.port
):
return JSONResponse({"detail": "Ungültiger Ursprung"}, status_code=403)
response = await call_next(request)
response.headers["Cache-Control"] = "no-store"
response.headers["X-Content-Type-Options"] = "nosniff"
response.headers["Referrer-Policy"] = "no-referrer"
response.headers["Content-Security-Policy"] = (
"default-src 'self'; script-src 'self'; style-src 'self'; "
"connect-src 'self'; media-src 'self' blob:; img-src 'self' data:; frame-ancestors 'none'"
)
return response
@app.get("/")
async def index():
return FileResponse(ROOT / "web" / "index.html")
@app.get("/api/status")
async def health():
return {**status, "busy": bool(sessions), "model": config.model, "thinking": False}
@app.post("/api/offer")
async def offer(body: Offer):
if not status["ready"]:
raise HTTPException(503, status["error"] or status["stage"])
async with lock:
if sessions:
raise HTTPException(
409, "Ein Gespräch läuft bereits. Beende es im anderen Fenster."
)
from pipecat.transports.smallwebrtc.connection import SmallWebRTCConnection
connection = SmallWebRTCConnection(ice_servers=[], connection_timeout_secs=20)
try:
await connection.initialize(body.sdp, body.type)
except Exception:
await connection.disconnect()
raise
async def session():
try:
await run_voice(connection, models, config)
except asyncio.CancelledError:
raise
except Exception as exc:
logger.exception("Gespräch fehlgeschlagen")
connection.send_app_message({"type": "error", "message": str(exc)})
finally:
await connection.disconnect()
sessions.pop(connection.pc_id, None)
task = asyncio.create_task(session())
sessions[connection.pc_id] = (connection, task)
return connection.get_answer()
@app.post("/api/disconnect")
async def disconnect(body: Disconnect):
pair = sessions.get(body.pc_id)
if pair:
connection, task = pair
await connection.disconnect()
task.cancel()
with contextlib.suppress(asyncio.CancelledError):
await task
return {"ok": True}
app.mount("/static", StaticFiles(directory=ROOT / "web"), name="static")
return app
def main():
os.environ.setdefault("HF_HUB_OFFLINE", "1")
logger.remove()
logger.add(sys.stderr, level=os.getenv("ASTRA_LOG_LEVEL", "INFO"))
config = Settings.from_env()
uvicorn.run(create_app(config), host="127.0.0.1", port=config.port, access_log=False)
if __name__ == "__main__":
main()