385 lines
13 KiB
Python
385 lines
13 KiB
Python
"""Deterministic WebSocket contract tests for the voice gateway route."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import unittest
|
|
from types import SimpleNamespace
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
from .deps import Principal, Role
|
|
from .persona_repository import PersonaVoiceMap
|
|
from .routes import voice as voice_routes
|
|
from .services.voice import VoicePreset
|
|
|
|
|
|
SESSION_ID = "voice-ws-contract-session"
|
|
VOICE_PRESET = VoicePreset(preset="neutral", openai_voice="sage")
|
|
|
|
|
|
def _principal(role: Role = Role.LEARNER) -> Principal:
|
|
return Principal(
|
|
user_id="00000000-0000-0000-0000-000000000201",
|
|
role=role,
|
|
cohort_ids=[],
|
|
email=f"voice-ws-{role.value}@hs.ac.kr",
|
|
display_name="Voice WS Contract",
|
|
)
|
|
|
|
|
|
def _control(payload: dict[str, object]) -> dict[str, object]:
|
|
return {"text": json.dumps(payload)}
|
|
|
|
|
|
def _binary(data: bytes) -> dict[str, object]:
|
|
return {"bytes": data}
|
|
|
|
|
|
class FakeWebSocket:
|
|
def __init__(self, incoming: list[dict[str, object]] | None = None) -> None:
|
|
self._incoming = list(incoming or [])
|
|
self.accepted = False
|
|
self.client_state = voice_routes.WebSocketState.CONNECTING
|
|
self.cookies: dict[str, str] = {}
|
|
self.query_params: dict[str, str] = {}
|
|
self.sent_json: list[dict[str, object]] = []
|
|
self.sent_text: list[str] = []
|
|
self.sent_bytes: list[bytes] = []
|
|
self.close_codes: list[int] = []
|
|
|
|
async def accept(self) -> None:
|
|
self.accepted = True
|
|
self.client_state = voice_routes.WebSocketState.CONNECTED
|
|
|
|
async def receive(self) -> dict[str, object]:
|
|
if self._incoming:
|
|
return self._incoming.pop(0)
|
|
return {"type": "websocket.disconnect"}
|
|
|
|
async def send_text(self, data: str) -> None:
|
|
self.sent_text.append(data)
|
|
self.sent_json.append(json.loads(data))
|
|
|
|
async def send_bytes(self, data: bytes) -> None:
|
|
self.sent_bytes.append(data)
|
|
|
|
async def close(self, code: int = 1000) -> None:
|
|
self.close_codes.append(code)
|
|
self.client_state = voice_routes.WebSocketState.DISCONNECTED
|
|
|
|
|
|
class VoiceWebSocketContractTest(unittest.IsolatedAsyncioTestCase):
|
|
def _bind_result(self) -> tuple[str, VoicePreset, None, dict[str, object]]:
|
|
return (
|
|
SESSION_ID,
|
|
VOICE_PRESET,
|
|
None,
|
|
{"degraded": False, "persona_catalog_source": "session"},
|
|
)
|
|
|
|
async def test_audio_start_binary_chunks_audio_end_ping_close_contract(self) -> None:
|
|
websocket = FakeWebSocket(
|
|
[
|
|
_control({"type": "audio_start", "format": "webm"}),
|
|
_binary(b"chunk-one"),
|
|
_binary(b"chunk-two"),
|
|
_control({"type": "ping"}),
|
|
_control(
|
|
{
|
|
"type": "audio_end",
|
|
"format": "webm",
|
|
"silence_ms": "450",
|
|
"barge_in": "true",
|
|
}
|
|
),
|
|
_control({"type": "close"}),
|
|
]
|
|
)
|
|
handle_utterance = AsyncMock()
|
|
|
|
with patch.object(
|
|
voice_routes,
|
|
"_principal_from_websocket",
|
|
AsyncMock(return_value=_principal()),
|
|
), patch.object(
|
|
voice_routes,
|
|
"_bind_session",
|
|
AsyncMock(return_value=self._bind_result()),
|
|
), patch.object(
|
|
voice_routes.voice_service,
|
|
"is_available",
|
|
return_value=True,
|
|
), patch.object(
|
|
voice_routes,
|
|
"_handle_utterance",
|
|
handle_utterance,
|
|
), patch.object(
|
|
voice_routes,
|
|
"_run_turn_and_speak",
|
|
AsyncMock(),
|
|
), patch.object(
|
|
voice_routes.time,
|
|
"monotonic",
|
|
side_effect=[10.0, 12.0],
|
|
):
|
|
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
|
|
|
self.assertTrue(websocket.accepted)
|
|
self.assertEqual(websocket.close_codes, [1000])
|
|
self.assertEqual(
|
|
[(message.get("type"), message.get("state")) for message in websocket.sent_json],
|
|
[("ready", "idle"), ("state", "listening"), ("pong", None)],
|
|
)
|
|
handle_utterance.assert_awaited_once()
|
|
kwargs = handle_utterance.await_args.kwargs
|
|
self.assertEqual(kwargs["session_id"], SESSION_ID)
|
|
self.assertEqual(kwargs["principal"].user_id, _principal().user_id)
|
|
self.assertEqual(kwargs["voice_preset"], VOICE_PRESET)
|
|
self.assertEqual(kwargs["audio"], b"chunk-onechunk-two")
|
|
self.assertEqual(kwargs["fmt"], "webm")
|
|
self.assertEqual(kwargs["audio_started_at"], 10.0)
|
|
self.assertEqual(kwargs["audio_ended_at"], 12.0)
|
|
self.assertEqual(kwargs["silence_ms"], 450)
|
|
self.assertIs(kwargs["barge_in"], True)
|
|
|
|
async def test_text_turn_strips_text_runs_turn_and_ping_close_still_work(self) -> None:
|
|
websocket = FakeWebSocket(
|
|
[
|
|
_control({"type": "text_turn", "text": " I need help practicing. "}),
|
|
_control({"type": "ping"}),
|
|
_control({"type": "close"}),
|
|
]
|
|
)
|
|
run_turn = AsyncMock()
|
|
handle_utterance = AsyncMock()
|
|
|
|
with patch.object(
|
|
voice_routes,
|
|
"_principal_from_websocket",
|
|
AsyncMock(return_value=_principal()),
|
|
), patch.object(
|
|
voice_routes,
|
|
"_bind_session",
|
|
AsyncMock(return_value=self._bind_result()),
|
|
), patch.object(
|
|
voice_routes.voice_service,
|
|
"is_available",
|
|
return_value=True,
|
|
), patch.object(
|
|
voice_routes,
|
|
"_handle_utterance",
|
|
handle_utterance,
|
|
), patch.object(
|
|
voice_routes,
|
|
"_run_turn_and_speak",
|
|
run_turn,
|
|
):
|
|
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
|
|
|
self.assertEqual(websocket.close_codes, [1000])
|
|
self.assertEqual(
|
|
[(message.get("type"), message.get("state")) for message in websocket.sent_json],
|
|
[("ready", "idle"), ("pong", None)],
|
|
)
|
|
handle_utterance.assert_not_awaited()
|
|
run_turn.assert_awaited_once()
|
|
kwargs = run_turn.await_args.kwargs
|
|
self.assertEqual(kwargs["session_id"], SESSION_ID)
|
|
self.assertEqual(kwargs["principal"].user_id, _principal().user_id)
|
|
self.assertEqual(kwargs["voice_preset"], VOICE_PRESET)
|
|
self.assertEqual(kwargs["learner_text"], "I need help practicing.")
|
|
|
|
async def test_oversize_binary_audio_reports_error_and_drops_utterance(self) -> None:
|
|
websocket = FakeWebSocket(
|
|
[
|
|
_control({"type": "audio_start", "format": "webm"}),
|
|
_binary(b"12345"),
|
|
_control({"type": "close"}),
|
|
]
|
|
)
|
|
handle_utterance = AsyncMock()
|
|
run_turn = AsyncMock()
|
|
|
|
with patch.object(
|
|
voice_routes,
|
|
"_principal_from_websocket",
|
|
AsyncMock(return_value=_principal()),
|
|
), patch.object(
|
|
voice_routes,
|
|
"_bind_session",
|
|
AsyncMock(return_value=self._bind_result()),
|
|
), patch.object(
|
|
voice_routes.voice_service,
|
|
"is_available",
|
|
return_value=True,
|
|
), patch.object(
|
|
voice_routes,
|
|
"_handle_utterance",
|
|
handle_utterance,
|
|
), patch.object(
|
|
voice_routes,
|
|
"_run_turn_and_speak",
|
|
run_turn,
|
|
), patch.object(
|
|
voice_routes,
|
|
"_MAX_AUDIO_BYTES",
|
|
4,
|
|
):
|
|
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
|
|
|
self.assertEqual(websocket.close_codes, [1000])
|
|
self.assertEqual(
|
|
websocket.sent_json[-1],
|
|
{
|
|
"type": "error",
|
|
"detail": "audio too large; please send a shorter utterance",
|
|
},
|
|
)
|
|
handle_utterance.assert_not_awaited()
|
|
run_turn.assert_not_awaited()
|
|
|
|
async def test_unauthenticated_client_closes_before_session_or_voice_checks(self) -> None:
|
|
websocket = FakeWebSocket()
|
|
bind_session = AsyncMock()
|
|
|
|
with patch.object(
|
|
voice_routes,
|
|
"_principal_from_websocket",
|
|
AsyncMock(return_value=None),
|
|
), patch.object(
|
|
voice_routes,
|
|
"_bind_session",
|
|
bind_session,
|
|
), patch.object(
|
|
voice_routes.voice_service,
|
|
"is_available",
|
|
return_value=True,
|
|
) as is_available:
|
|
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
|
|
|
self.assertTrue(websocket.accepted)
|
|
self.assertEqual(websocket.close_codes, [voice_routes.WS_CLOSE_UNAUTHORIZED])
|
|
self.assertEqual(websocket.sent_json, [{"type": "error", "detail": "not authenticated"}])
|
|
bind_session.assert_not_awaited()
|
|
is_available.assert_not_called()
|
|
|
|
async def test_non_learner_client_closes_before_session_or_voice_checks(self) -> None:
|
|
websocket = FakeWebSocket()
|
|
bind_session = AsyncMock()
|
|
|
|
with patch.object(
|
|
voice_routes,
|
|
"_principal_from_websocket",
|
|
AsyncMock(return_value=_principal(Role.TEACHER)),
|
|
), patch.object(
|
|
voice_routes,
|
|
"_bind_session",
|
|
bind_session,
|
|
), patch.object(
|
|
voice_routes.voice_service,
|
|
"is_available",
|
|
return_value=True,
|
|
) as is_available:
|
|
await voice_routes.voice_ws(websocket) # type: ignore[arg-type]
|
|
|
|
self.assertTrue(websocket.accepted)
|
|
self.assertEqual(websocket.close_codes, [voice_routes.WS_CLOSE_UNAUTHORIZED])
|
|
self.assertEqual(
|
|
websocket.sent_json,
|
|
[{"type": "error", "detail": "only learners can use voice"}],
|
|
)
|
|
bind_session.assert_not_awaited()
|
|
is_available.assert_not_called()
|
|
|
|
async def test_bind_session_uses_db_voice_map_for_existing_session(self) -> None:
|
|
websocket = FakeWebSocket()
|
|
websocket.query_params = {"session_id": SESSION_ID}
|
|
sess = SimpleNamespace(persona=SimpleNamespace(code="P2"))
|
|
voice_map = PersonaVoiceMap(
|
|
provider="openai",
|
|
voice_id="voice-p2-custom",
|
|
base_params={
|
|
"preset": "calm-adult-male",
|
|
"openai_voice": "onyx",
|
|
"rate": 1.08,
|
|
"instructions": "Low, guarded delivery.",
|
|
},
|
|
)
|
|
get_voice_map = AsyncMock(return_value=voice_map)
|
|
|
|
with patch.object(
|
|
voice_routes,
|
|
"_load_voice_session",
|
|
AsyncMock(return_value=(sess, None)),
|
|
), patch.object(
|
|
voice_routes,
|
|
"get_session_voice_map",
|
|
get_voice_map,
|
|
):
|
|
session_id, voice, err, meta = await voice_routes._bind_session(
|
|
websocket, _principal()
|
|
)
|
|
|
|
self.assertEqual(session_id, SESSION_ID)
|
|
self.assertIsNone(err)
|
|
self.assertEqual(meta["persona_catalog_source"], "session")
|
|
self.assertEqual(voice, VoicePreset(
|
|
preset="calm-adult-male",
|
|
openai_voice="onyx",
|
|
rate=1.08,
|
|
instructions="Low, guarded delivery.",
|
|
))
|
|
get_voice_map.assert_awaited_once_with(SESSION_ID)
|
|
|
|
async def test_bind_session_explicit_preset_overrides_db_voice_map(self) -> None:
|
|
websocket = FakeWebSocket()
|
|
websocket.query_params = {"session_id": SESSION_ID, "preset": "soft-young-fem"}
|
|
sess = SimpleNamespace(persona=SimpleNamespace(code="P2"))
|
|
get_voice_map = AsyncMock()
|
|
|
|
with patch.object(
|
|
voice_routes,
|
|
"_load_voice_session",
|
|
AsyncMock(return_value=(sess, None)),
|
|
), patch.object(
|
|
voice_routes,
|
|
"get_session_voice_map",
|
|
get_voice_map,
|
|
):
|
|
session_id, voice, err, _ = await voice_routes._bind_session(
|
|
websocket, _principal()
|
|
)
|
|
|
|
self.assertEqual(session_id, SESSION_ID)
|
|
self.assertIsNone(err)
|
|
self.assertEqual(voice.preset, "soft-young-fem")
|
|
self.assertEqual(voice.openai_voice, "coral")
|
|
get_voice_map.assert_not_awaited()
|
|
|
|
async def test_catalog_voice_map_is_used_for_dev_persona_binding_helper(self) -> None:
|
|
voice_map = PersonaVoiceMap(
|
|
provider="openai",
|
|
voice_id="verse",
|
|
base_params={"preset": "soft-young-fem", "rate": 0.9},
|
|
)
|
|
get_voice_map = AsyncMock(return_value=voice_map)
|
|
|
|
with patch.object(voice_routes, "get_persona_voice_map", get_voice_map):
|
|
voice = await voice_routes._resolve_catalog_voice(
|
|
persona_id="00000000-0000-0000-0000-000000000301",
|
|
version=7,
|
|
persona_code="P1",
|
|
explicit_preset=None,
|
|
)
|
|
|
|
self.assertEqual(voice.preset, "soft-young-fem")
|
|
self.assertEqual(voice.openai_voice, "verse")
|
|
self.assertAlmostEqual(voice.rate, 0.9)
|
|
get_voice_map.assert_awaited_once_with(
|
|
persona_id="00000000-0000-0000-0000-000000000301",
|
|
version=7,
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|