独立依赖
This commit is contained in:
@@ -4,15 +4,13 @@ import asyncio
|
||||
import io
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import wave
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import numpy as np
|
||||
from aiortc import RTCPeerConnection, RTCSessionDescription
|
||||
from dotenv import load_dotenv
|
||||
from fastapi import FastAPI, Request, WebSocket, WebSocketDisconnect
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse, JSONResponse, StreamingResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
@@ -26,11 +24,8 @@ from services.avatar import AvatarService
|
||||
from services.llm import LLMService
|
||||
from services.tts import TTSService
|
||||
from services.vad import VADService
|
||||
from webrtc.tracks import AudioBus, AvatarAudioTrack
|
||||
|
||||
load_dotenv()
|
||||
if settings.hf_token:
|
||||
os.environ["HF_TOKEN"] = settings.hf_token
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
app = FastAPI(title="Visual Voice Chat")
|
||||
@@ -50,8 +45,6 @@ tts_service = TTSService()
|
||||
vad_service = VADService()
|
||||
avatar_service = AvatarService()
|
||||
pipeline = ChatPipeline(arbitrator, asr_service, llm_service, tts_service)
|
||||
audio_bus = AudioBus(sample_rate=48000)
|
||||
pcs: set[RTCPeerConnection] = set()
|
||||
subtitle_clients: set[WebSocket] = set()
|
||||
animation_clients: set[WebSocket] = set()
|
||||
audio_clients: set[WebSocket] = set()
|
||||
@@ -146,7 +139,6 @@ async def _broadcast_animation(payload: dict) -> None:
|
||||
def _runtime_payload() -> dict:
|
||||
return {
|
||||
"state": str(arbitrator.state),
|
||||
"peers": len(pcs),
|
||||
"audio_clients": len(audio_clients),
|
||||
"subtitle_clients": len(subtitle_clients),
|
||||
"animation_clients": len(animation_clients),
|
||||
@@ -341,7 +333,6 @@ async def _process_user_audio_chunk(mono_16k: np.ndarray, vad_buffer: np.ndarray
|
||||
barge_t0 = time.perf_counter()
|
||||
was_avatar_speaking = arbitrator.state == SessionState.AVATAR_SPEAKING
|
||||
await arbitrator.on_speech_start()
|
||||
await audio_bus.clear()
|
||||
await _broadcast_audio_reset("speech-start")
|
||||
await _broadcast_animation(await avatar_service.build_reset_payload(reason="speech-start"))
|
||||
if was_avatar_speaking:
|
||||
@@ -357,18 +348,6 @@ async def _process_user_audio_chunk(mono_16k: np.ndarray, vad_buffer: np.ndarray
|
||||
return vad_buffer
|
||||
|
||||
|
||||
async def _consume_user_audio(track) -> None:
|
||||
vad_buffer = np.zeros(0, dtype=np.float32)
|
||||
while True:
|
||||
frame = await track.recv()
|
||||
pcm = frame.to_ndarray()
|
||||
sample_rate = getattr(frame, "sample_rate", 48000) or 48000
|
||||
layout = getattr(frame, "layout", None)
|
||||
channels = len(layout.channels) if layout and layout.channels else 1
|
||||
mono_16k = _resample_to_16k_mono(pcm, sample_rate, channels)
|
||||
vad_buffer = await _process_user_audio_chunk(mono_16k, vad_buffer)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return _runtime_payload()
|
||||
@@ -377,10 +356,9 @@ async def health():
|
||||
@app.get("/meta")
|
||||
async def meta():
|
||||
return {
|
||||
"stun_url": settings.stun_url,
|
||||
"https_enabled": bool(settings.ssl_certfile and settings.ssl_keyfile),
|
||||
"host": settings.webrtc_host,
|
||||
"port": settings.webrtc_port,
|
||||
"host": settings.http_host,
|
||||
"port": settings.http_port,
|
||||
"avatar_protocol": settings.avatar_control_protocol,
|
||||
}
|
||||
|
||||
@@ -474,7 +452,6 @@ async def ws_audio(websocket: WebSocket):
|
||||
elif event_type == "audio_reset":
|
||||
_clear_buffered_audio()
|
||||
runtime.pending_voice_turn = False
|
||||
await audio_bus.clear()
|
||||
await _broadcast_audio_reset("client-reset")
|
||||
elif event_type == "ping":
|
||||
await websocket.send_text(json.dumps({"type": "pong"}, ensure_ascii=False))
|
||||
@@ -567,7 +544,6 @@ async def chat_reset():
|
||||
runtime.vad_start_count = 0
|
||||
runtime.vad_end_count = 0
|
||||
runtime.last_animation_frame_count = 0
|
||||
await audio_bus.clear()
|
||||
await _broadcast_audio_reset("chat-reset")
|
||||
await _broadcast_animation(await avatar_service.build_reset_payload(reason="chat-reset"))
|
||||
return {"ok": True}
|
||||
@@ -583,60 +559,13 @@ async def root():
|
||||
return FileResponse("web/index.html")
|
||||
|
||||
|
||||
@app.post("/webrtc/offer")
|
||||
async def webrtc_offer(request: Request):
|
||||
params = await request.json()
|
||||
offer = RTCSessionDescription(sdp=params["sdp"], type=params["type"])
|
||||
|
||||
replaced_previous = len(pcs) > 0
|
||||
if replaced_previous:
|
||||
old_peers = list(pcs)
|
||||
await asyncio.gather(*(peer.close() for peer in old_peers), return_exceptions=True)
|
||||
for peer in old_peers:
|
||||
pcs.discard(peer)
|
||||
|
||||
pc = RTCPeerConnection()
|
||||
pcs.add(pc)
|
||||
|
||||
@pc.on("connectionstatechange")
|
||||
async def on_connectionstatechange():
|
||||
if pc.connectionState in {"failed", "closed", "disconnected"}:
|
||||
await pc.close()
|
||||
pcs.discard(pc)
|
||||
|
||||
@pc.on("track")
|
||||
def on_track(track):
|
||||
if track.kind == "audio":
|
||||
asyncio.create_task(_consume_user_audio(track))
|
||||
|
||||
pc.addTrack(AvatarAudioTrack(audio_bus=audio_bus))
|
||||
|
||||
await pc.setRemoteDescription(offer)
|
||||
answer = await pc.createAnswer()
|
||||
await pc.setLocalDescription(answer)
|
||||
|
||||
return JSONResponse(
|
||||
{
|
||||
"sdp": pc.localDescription.sdp,
|
||||
"type": pc.localDescription.type,
|
||||
"replaced_previous": replaced_previous,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@app.on_event("shutdown")
|
||||
async def on_shutdown():
|
||||
await asyncio.gather(*(pc.close() for pc in list(pcs)), return_exceptions=True)
|
||||
pcs.clear()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import uvicorn
|
||||
|
||||
uvicorn_kwargs = {
|
||||
"app": app,
|
||||
"host": settings.webrtc_host,
|
||||
"port": settings.webrtc_port,
|
||||
"host": settings.http_host,
|
||||
"port": settings.http_port,
|
||||
}
|
||||
if settings.ssl_certfile and settings.ssl_keyfile:
|
||||
uvicorn_kwargs["ssl_certfile"] = settings.ssl_certfile
|
||||
|
||||
Reference in New Issue
Block a user