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:
Jamie Pine
2026-03-16 23:17:09 -07:00
parent 01800f196f
commit f9e1aa153d
2 changed files with 30 additions and 19 deletions
+9
View File
@@ -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(),
+2
View File
@@ -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: