mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 00:55:14 -07:00
handle client disconnects in SSE and streaming endpoints
Wrap SSE generators with BrokenPipeError/ConnectionResetError handling so client disconnects during generation status polling, download progress, or audio streaming don't produce unhandled Errno 32 errors. Closes #248
This commit is contained in:
@@ -1,12 +1,15 @@
|
|||||||
"""TTS generation endpoints."""
|
"""TTS generation endpoints."""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException
|
from fastapi import APIRouter, Depends, HTTPException
|
||||||
from fastapi.responses import StreamingResponse
|
from fastapi.responses import StreamingResponse
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
from .. import models
|
from .. import models
|
||||||
from ..services import history, profiles, tts
|
from ..services import history, profiles, tts
|
||||||
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
|
from ..database import Generation as DBGeneration, VoiceProfile as DBVoiceProfile, get_db
|
||||||
@@ -181,6 +184,7 @@ async def get_generation_status(generation_id: str, db: Session = Depends(get_db
|
|||||||
import json
|
import json
|
||||||
|
|
||||||
async def event_stream():
|
async def event_stream():
|
||||||
|
try:
|
||||||
while True:
|
while True:
|
||||||
db.expire_all()
|
db.expire_all()
|
||||||
gen = db.query(DBGeneration).filter_by(id=generation_id).first()
|
gen = db.query(DBGeneration).filter_by(id=generation_id).first()
|
||||||
@@ -200,6 +204,8 @@ async def get_generation_status(generation_id: str, db: Session = Depends(get_db
|
|||||||
return
|
return
|
||||||
|
|
||||||
await asyncio.sleep(1)
|
await asyncio.sleep(1)
|
||||||
|
except (BrokenPipeError, ConnectionResetError, asyncio.CancelledError):
|
||||||
|
logger.debug("SSE client disconnected for generation %s", generation_id)
|
||||||
|
|
||||||
return StreamingResponse(
|
return StreamingResponse(
|
||||||
event_stream(),
|
event_stream(),
|
||||||
@@ -265,9 +271,12 @@ async def stream_speech(
|
|||||||
wav_bytes = tts.audio_to_wav_bytes(audio, sample_rate)
|
wav_bytes = tts.audio_to_wav_bytes(audio, sample_rate)
|
||||||
|
|
||||||
async def _wav_stream():
|
async def _wav_stream():
|
||||||
|
try:
|
||||||
chunk_size = 64 * 1024
|
chunk_size = 64 * 1024
|
||||||
for i in range(0, len(wav_bytes), chunk_size):
|
for i in range(0, len(wav_bytes), chunk_size):
|
||||||
yield wav_bytes[i : i + chunk_size]
|
yield wav_bytes[i : i + chunk_size]
|
||||||
|
except (BrokenPipeError, ConnectionResetError, asyncio.CancelledError):
|
||||||
|
logger.debug("Client disconnected during audio stream")
|
||||||
|
|
||||||
return StreamingResponse(
|
return StreamingResponse(
|
||||||
_wav_stream(),
|
_wav_stream(),
|
||||||
|
|||||||
@@ -246,6 +246,8 @@ class ProgressManager:
|
|||||||
# Send heartbeat
|
# Send heartbeat
|
||||||
yield ": heartbeat\n\n"
|
yield ": heartbeat\n\n"
|
||||||
continue
|
continue
|
||||||
|
except (BrokenPipeError, ConnectionResetError, asyncio.CancelledError):
|
||||||
|
logger.debug(f"SSE client disconnected from {model_name}")
|
||||||
finally:
|
finally:
|
||||||
# Remove from listeners
|
# Remove from listeners
|
||||||
if model_name in self._listeners:
|
if model_name in self._listeners:
|
||||||
|
|||||||
Reference in New Issue
Block a user