텍스트 응답 음성 재생 연결
This commit is contained in:
parent
64e06a1185
commit
d80e33da5e
9 changed files with 524 additions and 16 deletions
|
|
@ -18,14 +18,15 @@ import hashlib
|
|||
import time
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, WebSocket, WebSocketDisconnect, HTTPException
|
||||
from fastapi.responses import JSONResponse
|
||||
from fastapi import APIRouter, WebSocket, WebSocketDisconnect, HTTPException, status
|
||||
from fastapi.responses import JSONResponse, Response
|
||||
from pydantic import BaseModel, Field
|
||||
from starlette.websockets import WebSocketState
|
||||
|
||||
from .. import session_persistence, turn_runtime
|
||||
from ..auth_sessions import get_session, user_has_consent, user_onboarding_complete
|
||||
from ..config import settings
|
||||
from ..deps import Principal, Role
|
||||
from ..deps import CurrentPrincipal, Principal, Role
|
||||
from ..engine_client import EngineError, engine_client
|
||||
from ..persona_repository import (
|
||||
PersonaVoiceMap,
|
||||
|
|
@ -41,6 +42,13 @@ from ..store import InProcSession, TurnRecord, store
|
|||
|
||||
router = APIRouter(prefix="/voice", tags=["voice"])
|
||||
|
||||
|
||||
class VoiceSpeechRequest(BaseModel):
|
||||
"""Request OpenAI TTS for an already-persisted client reply."""
|
||||
|
||||
session_id: str = Field(min_length=1, max_length=80)
|
||||
turn_seq: int = Field(ge=1)
|
||||
|
||||
# WebSocket close codes.
|
||||
WS_CLOSE_DEGRADED = 1011
|
||||
WS_CLOSE_BAD_REQUEST = 1008
|
||||
|
|
@ -123,6 +131,89 @@ async def voice_health() -> JSONResponse:
|
|||
return JSONResponse(body, status_code=200 if available else 503)
|
||||
|
||||
|
||||
@router.post("/speech")
|
||||
async def voice_speech(body: VoiceSpeechRequest, principal: CurrentPrincipal) -> Response:
|
||||
"""Synthesize the persisted client reply for a completed text turn.
|
||||
|
||||
The browser sends only session/turn identifiers. The server reloads the
|
||||
owned session and speaks the stored client-visible reply, so this endpoint
|
||||
cannot be used as an arbitrary paid text-to-speech proxy.
|
||||
"""
|
||||
learner = principal
|
||||
if learner.role != Role.LEARNER and learner.can_access_role(Role.LEARNER):
|
||||
learner = learner.with_role(Role.LEARNER)
|
||||
if learner.role != Role.LEARNER:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN,
|
||||
detail="only learners can use voice",
|
||||
)
|
||||
|
||||
access_error = await _practice_access_error(learner)
|
||||
if access_error is not None:
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail=access_error)
|
||||
if not voice_service.is_available():
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="OPENAI_API_KEY is not configured",
|
||||
)
|
||||
|
||||
sess, err = await turn_runtime.load_owned_session(
|
||||
body.session_id,
|
||||
learner,
|
||||
allow_ended=False,
|
||||
)
|
||||
if sess is None:
|
||||
status_code = {
|
||||
turn_runtime.SessionAccessError.FORBIDDEN: status.HTTP_403_FORBIDDEN,
|
||||
turn_runtime.SessionAccessError.ENDED: status.HTTP_409_CONFLICT,
|
||||
}.get(err, status.HTTP_404_NOT_FOUND)
|
||||
raise HTTPException(status_code=status_code, detail=f"voice session {err or 'not_found'}")
|
||||
|
||||
text = _client_turn_text_for_speech(sess, body.turn_seq)
|
||||
if text is None:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="client reply not found for turn",
|
||||
)
|
||||
|
||||
voice_preset = await _resolve_session_voice(
|
||||
session_id=body.session_id,
|
||||
persona_code=sess.persona.code,
|
||||
explicit_preset=None,
|
||||
)
|
||||
try:
|
||||
chunks = [
|
||||
chunk.audio
|
||||
async for chunk in voice_service.synthesize_stream(text, voice_preset)
|
||||
]
|
||||
except VoiceUnavailable as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=str(exc),
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
||||
detail=f"TTS failed: {exc}",
|
||||
) from exc
|
||||
|
||||
audio = b"".join(chunks)
|
||||
if not audio:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||
detail="client reply has no speakable text",
|
||||
)
|
||||
return Response(
|
||||
content=audio,
|
||||
media_type="audio/mpeg",
|
||||
headers={
|
||||
"Cache-Control": "no-store",
|
||||
"X-Vignette-TTS-Model": voice_svc.TTS_MODEL,
|
||||
"X-Vignette-TTS-Provider": voice_service.tts_provider(),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@router.websocket("/ws")
|
||||
async def voice_ws(websocket: WebSocket) -> None:
|
||||
"""Run one authenticated learner voice cascade."""
|
||||
|
|
@ -596,6 +687,19 @@ async def _load_voice_session(
|
|||
return sess, None
|
||||
|
||||
|
||||
def _client_turn_text_for_speech(sess: InProcSession, turn_seq: int) -> str | None:
|
||||
"""Return the persisted client-visible reply for one completed turn."""
|
||||
for turn in reversed(sess.turns):
|
||||
if (
|
||||
turn.turn_seq == turn_seq
|
||||
and turn.speaker == "client"
|
||||
and turn.is_visible_to("client")
|
||||
):
|
||||
text = (turn.text_masked or turn.text).strip()
|
||||
return text or None
|
||||
return None
|
||||
|
||||
|
||||
async def _principal_from_websocket(websocket: WebSocket) -> Principal | None:
|
||||
"""Restore the same server-side browser session used by REST routes."""
|
||||
raw_cookie = websocket.cookies.get(settings.cookie_name)
|
||||
|
|
|
|||
|
|
@ -12,7 +12,8 @@ from fastapi import HTTPException
|
|||
from .deps import Principal, Role
|
||||
from .persona_repository import PersonaVoiceMap
|
||||
from .routes import voice as voice_routes
|
||||
from .services.voice import VoicePreset
|
||||
from .services.voice import TTSChunk, VoicePreset
|
||||
from .store import TurnRecord
|
||||
|
||||
|
||||
SESSION_ID = "voice-ws-contract-session"
|
||||
|
|
@ -130,6 +131,98 @@ class VoiceWebSocketContractTest(unittest.IsolatedAsyncioTestCase):
|
|||
],
|
||||
)
|
||||
|
||||
def test_text_tts_uses_only_the_persisted_client_visible_reply(self) -> None:
|
||||
session = SimpleNamespace(
|
||||
turns=[
|
||||
TurnRecord(
|
||||
turn_seq=2,
|
||||
speaker="counselor",
|
||||
stage="초기",
|
||||
text="raw learner text",
|
||||
text_masked="masked learner text",
|
||||
),
|
||||
TurnRecord(
|
||||
turn_seq=2,
|
||||
speaker="client",
|
||||
stage="초기",
|
||||
text="raw client reply",
|
||||
text_masked="마스킹된 내담자 응답",
|
||||
),
|
||||
TurnRecord(
|
||||
turn_seq=3,
|
||||
speaker="client",
|
||||
stage="초기",
|
||||
text="hidden evaluator reply",
|
||||
text_masked="hidden evaluator reply",
|
||||
visible_to=("evaluator",),
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
voice_routes._client_turn_text_for_speech(session, 2),
|
||||
"마스킹된 내담자 응답",
|
||||
)
|
||||
self.assertIsNone(voice_routes._client_turn_text_for_speech(session, 3))
|
||||
self.assertIsNone(voice_routes._client_turn_text_for_speech(session, 99))
|
||||
|
||||
async def test_text_turn_speech_returns_openai_audio_for_owned_persisted_turn(self) -> None:
|
||||
session = SimpleNamespace(
|
||||
persona=SimpleNamespace(code="P1"),
|
||||
turns=[
|
||||
TurnRecord(
|
||||
turn_seq=4,
|
||||
speaker="client",
|
||||
stage="초기",
|
||||
text="내담자 응답",
|
||||
text_masked="내담자 응답",
|
||||
)
|
||||
],
|
||||
)
|
||||
synthesized: list[tuple[str, VoicePreset]] = []
|
||||
|
||||
async def synthesize(text: str, voice: VoicePreset):
|
||||
synthesized.append((text, voice))
|
||||
yield TTSChunk(audio=b"mp3-a")
|
||||
yield TTSChunk(audio=b"mp3-b")
|
||||
|
||||
with patch.object(
|
||||
voice_routes,
|
||||
"_practice_access_error",
|
||||
AsyncMock(return_value=None),
|
||||
), patch.object(
|
||||
voice_routes.turn_runtime,
|
||||
"load_owned_session",
|
||||
AsyncMock(return_value=(session, None)),
|
||||
), patch.object(
|
||||
voice_routes,
|
||||
"_resolve_session_voice",
|
||||
AsyncMock(return_value=VOICE_PRESET),
|
||||
), patch.object(
|
||||
voice_routes.voice_service,
|
||||
"is_available",
|
||||
return_value=True,
|
||||
), patch.object(
|
||||
voice_routes.voice_service,
|
||||
"tts_provider",
|
||||
return_value="openai",
|
||||
), patch.object(
|
||||
voice_routes.voice_service,
|
||||
"synthesize_stream",
|
||||
new=synthesize,
|
||||
):
|
||||
response = await voice_routes.voice_speech(
|
||||
voice_routes.VoiceSpeechRequest(session_id=SESSION_ID, turn_seq=4),
|
||||
_principal(),
|
||||
)
|
||||
|
||||
self.assertEqual(response.status_code, 200)
|
||||
self.assertEqual(response.media_type, "audio/mpeg")
|
||||
self.assertEqual(response.body, b"mp3-amp3-b")
|
||||
self.assertEqual(response.headers["cache-control"], "no-store")
|
||||
self.assertEqual(response.headers["x-vignette-tts-provider"], "openai")
|
||||
self.assertEqual(synthesized, [("내담자 응답", VOICE_PRESET)])
|
||||
|
||||
async def test_audio_start_binary_chunks_audio_end_ping_close_contract(self) -> None:
|
||||
websocket = FakeWebSocket(
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue