런타임 계약과 학습자 흐름 보강
This commit is contained in:
parent
f456b8997a
commit
206018b088
56 changed files with 4306 additions and 1008 deletions
|
|
@ -177,6 +177,10 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
receiving = False
|
||||
audio_started_at: float | None = None
|
||||
last_audio_end_at: float | None = None
|
||||
audio_format: str | None = None
|
||||
audio_sample_rate: int | None = None
|
||||
audio_channels: int | None = None
|
||||
audio_sample_width: int | None = None
|
||||
|
||||
try:
|
||||
while True:
|
||||
|
|
@ -217,6 +221,10 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
if ctype == "audio_start":
|
||||
receiving = True
|
||||
audio_started_at = time.monotonic()
|
||||
audio_format = _safe_str(ctrl.get("format"))
|
||||
audio_sample_rate = _safe_int(ctrl.get("sample_rate"))
|
||||
audio_channels = _safe_int(ctrl.get("channels"))
|
||||
audio_sample_width = _safe_int(ctrl.get("sample_width"))
|
||||
audio_buf.clear()
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "listening"})
|
||||
|
||||
|
|
@ -226,13 +234,17 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
silence_ms = _safe_int(ctrl.get("silence_ms"))
|
||||
if silence_ms is None and last_audio_end_at is not None and audio_started_at is not None:
|
||||
silence_ms = max(0, int((audio_started_at - last_audio_end_at) * 1000))
|
||||
end_format = _safe_str(ctrl.get("format")) or audio_format
|
||||
await _handle_utterance(
|
||||
websocket,
|
||||
session_id=session_id,
|
||||
principal=principal,
|
||||
voice_preset=voice_preset,
|
||||
audio=bytes(audio_buf),
|
||||
fmt=ctrl.get("format"),
|
||||
fmt=end_format,
|
||||
sample_rate=_safe_int(ctrl.get("sample_rate")) or audio_sample_rate,
|
||||
channels=_safe_int(ctrl.get("channels")) or audio_channels,
|
||||
sample_width=_safe_int(ctrl.get("sample_width")) or audio_sample_width,
|
||||
audio_started_at=audio_started_at,
|
||||
audio_ended_at=audio_ended_at,
|
||||
silence_ms=silence_ms,
|
||||
|
|
@ -241,6 +253,10 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
)
|
||||
last_audio_end_at = audio_ended_at
|
||||
audio_started_at = None
|
||||
audio_format = None
|
||||
audio_sample_rate = None
|
||||
audio_channels = None
|
||||
audio_sample_width = None
|
||||
audio_buf.clear()
|
||||
|
||||
elif ctype == "text_turn":
|
||||
|
|
@ -257,6 +273,27 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
learner_text=learner_text,
|
||||
)
|
||||
|
||||
elif ctype == "stt_result":
|
||||
receiving = False
|
||||
audio_buf.clear()
|
||||
stt_received_at = time.monotonic()
|
||||
await _handle_stt_result_control(
|
||||
websocket,
|
||||
session_id=session_id,
|
||||
principal=principal,
|
||||
voice_preset=voice_preset,
|
||||
ctrl=ctrl,
|
||||
audio_started_at=audio_started_at,
|
||||
audio_ended_at=stt_received_at,
|
||||
last_audio_end_at=last_audio_end_at,
|
||||
)
|
||||
last_audio_end_at = stt_received_at
|
||||
audio_started_at = None
|
||||
audio_format = None
|
||||
audio_sample_rate = None
|
||||
audio_channels = None
|
||||
audio_sample_width = None
|
||||
|
||||
elif ctype == "ping":
|
||||
await _safe_send_json(websocket, {"type": "pong"})
|
||||
|
||||
|
|
@ -271,6 +308,60 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
await _safe_close(websocket)
|
||||
|
||||
|
||||
async def _handle_stt_result_control(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
session_id: str,
|
||||
principal: Principal,
|
||||
voice_preset: VoicePreset,
|
||||
ctrl: dict[str, object],
|
||||
audio_started_at: float | None = None,
|
||||
audio_ended_at: float | None = None,
|
||||
last_audio_end_at: float | None = None,
|
||||
) -> None:
|
||||
learner_text = str(ctrl.get("text") or "").strip()
|
||||
transcript_final = _safe_bool(ctrl.get("final"))
|
||||
silence_ms = _safe_int(ctrl.get("silence_ms"))
|
||||
if silence_ms is None and last_audio_end_at is not None and audio_started_at is not None:
|
||||
silence_ms = max(0, int((audio_started_at - last_audio_end_at) * 1000))
|
||||
provider_events = _safe_provider_events(ctrl.get("provider_events"))
|
||||
decision = voice_svc.assess_end_of_turn(
|
||||
transcript_text=learner_text,
|
||||
transcript_final=bool(transcript_final),
|
||||
silence_ms=silence_ms,
|
||||
)
|
||||
await _safe_send_json(
|
||||
websocket,
|
||||
{
|
||||
"type": "eot",
|
||||
"ready": decision.ready,
|
||||
"reason": decision.reason,
|
||||
"silence_ms": decision.silence_ms,
|
||||
"threshold_ms": decision.threshold_ms,
|
||||
},
|
||||
)
|
||||
if not decision.ready:
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "listening"})
|
||||
return
|
||||
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "thinking"})
|
||||
await _safe_send_json(
|
||||
websocket,
|
||||
{"type": "transcript", "text": learner_text, "final": True, "speaker": "counselor"},
|
||||
)
|
||||
await _run_turn_and_speak(
|
||||
websocket,
|
||||
session_id=session_id,
|
||||
principal=principal,
|
||||
voice_preset=voice_preset,
|
||||
learner_text=learner_text,
|
||||
duration_s=_elapsed_seconds(audio_started_at, audio_ended_at),
|
||||
silence_ms=decision.silence_ms,
|
||||
barge_in=_safe_bool(ctrl.get("barge_in")),
|
||||
provider_events=provider_events,
|
||||
)
|
||||
|
||||
|
||||
async def _handle_utterance(
|
||||
websocket: WebSocket,
|
||||
*,
|
||||
|
|
@ -279,6 +370,9 @@ async def _handle_utterance(
|
|||
voice_preset: VoicePreset,
|
||||
audio: bytes,
|
||||
fmt: Optional[str],
|
||||
sample_rate: int | None = None,
|
||||
channels: int | None = None,
|
||||
sample_width: int | None = None,
|
||||
audio_started_at: float | None = None,
|
||||
audio_ended_at: float | None = None,
|
||||
silence_ms: int | None = None,
|
||||
|
|
@ -293,10 +387,17 @@ async def _handle_utterance(
|
|||
|
||||
# STT begins after the learner stops speaking.
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "thinking"})
|
||||
filename, content_type = _audio_meta(fmt)
|
||||
upload_audio, upload_fmt = _normalize_audio_upload(
|
||||
audio,
|
||||
fmt=fmt,
|
||||
sample_rate=sample_rate,
|
||||
channels=channels,
|
||||
sample_width=sample_width,
|
||||
)
|
||||
filename, content_type = _audio_meta(upload_fmt)
|
||||
try:
|
||||
stt = await voice_service.transcribe(
|
||||
audio, filename=filename, content_type=content_type
|
||||
upload_audio, filename=filename, content_type=content_type
|
||||
)
|
||||
except VoiceUnavailable as e:
|
||||
await _safe_send_json(websocket, {"type": "degraded", "reason": str(e)})
|
||||
|
|
@ -308,7 +409,7 @@ async def _handle_utterance(
|
|||
return
|
||||
|
||||
learner_text = stt.text
|
||||
audio_ref = _voice_audio_ref(audio, fmt)
|
||||
audio_ref = _voice_audio_ref(upload_audio, upload_fmt)
|
||||
duration_s = stt.duration or _elapsed_seconds(audio_started_at, audio_ended_at)
|
||||
speech_rate = _estimate_speech_rate(learner_text, duration_s)
|
||||
provider_events = _merge_provider_events(provider_events, getattr(stt, "provider_events", []))
|
||||
|
|
@ -342,12 +443,15 @@ async def _run_turn_and_speak(
|
|||
voice_preset: VoicePreset,
|
||||
learner_text: str,
|
||||
audio_ref: str | None = None,
|
||||
duration_s: float | None = None,
|
||||
silence_ms: int | None = None,
|
||||
speech_rate: float | None = None,
|
||||
barge_in: bool | None = None,
|
||||
provider_events: list[dict[str, object]] | None = None,
|
||||
) -> None:
|
||||
"""Run one counseling turn and stream synthesized client speech."""
|
||||
if speech_rate is None:
|
||||
speech_rate = _estimate_speech_rate(learner_text, duration_s)
|
||||
sess, err = await _load_voice_session(session_id, principal)
|
||||
if sess is None:
|
||||
await _safe_send_json(websocket, {"type": "error", "detail": err or "session not found or ended"})
|
||||
|
|
@ -364,10 +468,12 @@ async def _run_turn_and_speak(
|
|||
card=sess.persona,
|
||||
state=sess.state,
|
||||
learner_text=learner_text,
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
kb_behavior_cues=kb_cues,
|
||||
memory=orchestrator.TurnMemory(
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
kb_behavior_cues=kb_cues,
|
||||
),
|
||||
theory_mode=sess.theory_mode,
|
||||
)
|
||||
assert ctx.state_after is not None
|
||||
|
|
@ -656,6 +762,56 @@ def _audio_meta(fmt: Optional[str]) -> tuple[str, str]:
|
|||
return table.get(f, ("audio.webm", "audio/webm"))
|
||||
|
||||
|
||||
def _normalize_audio_upload(
|
||||
audio: bytes,
|
||||
*,
|
||||
fmt: Optional[str],
|
||||
sample_rate: int | None = None,
|
||||
channels: int | None = None,
|
||||
sample_width: int | None = None,
|
||||
) -> tuple[bytes, str]:
|
||||
f = (fmt or "webm").lower().lstrip(".") or "webm"
|
||||
if f != "pcm":
|
||||
return audio, f
|
||||
if sample_width not in (None, 2):
|
||||
raise ValueError("pcm sample_width must be 2 bytes")
|
||||
return _wav_from_pcm16(
|
||||
audio,
|
||||
sample_rate=_bounded_int(sample_rate, default=48000, minimum=8000, maximum=96000),
|
||||
channels=_bounded_int(channels, default=1, minimum=1, maximum=2),
|
||||
), "wav"
|
||||
|
||||
|
||||
def _bounded_int(value: int | None, *, default: int, minimum: int, maximum: int) -> int:
|
||||
if value is None:
|
||||
return default
|
||||
return min(maximum, max(minimum, value))
|
||||
|
||||
|
||||
def _wav_from_pcm16(pcm: bytes, *, sample_rate: int, channels: int) -> bytes:
|
||||
byte_rate = sample_rate * channels * 2
|
||||
block_align = channels * 2
|
||||
data_size = len(pcm)
|
||||
header = b"".join(
|
||||
[
|
||||
b"RIFF",
|
||||
(36 + data_size).to_bytes(4, "little"),
|
||||
b"WAVE",
|
||||
b"fmt ",
|
||||
(16).to_bytes(4, "little"),
|
||||
(1).to_bytes(2, "little"),
|
||||
channels.to_bytes(2, "little"),
|
||||
sample_rate.to_bytes(4, "little"),
|
||||
byte_rate.to_bytes(4, "little"),
|
||||
block_align.to_bytes(2, "little"),
|
||||
(16).to_bytes(2, "little"),
|
||||
b"data",
|
||||
data_size.to_bytes(4, "little"),
|
||||
]
|
||||
)
|
||||
return header + pcm
|
||||
|
||||
|
||||
def _voice_audio_ref(audio: bytes, fmt: Optional[str]) -> str | None:
|
||||
if not audio:
|
||||
return None
|
||||
|
|
@ -688,6 +844,13 @@ def _safe_int(value: object) -> int | None:
|
|||
return None
|
||||
|
||||
|
||||
def _safe_str(value: object) -> str | None:
|
||||
if isinstance(value, str):
|
||||
text = value.strip()
|
||||
return text or None
|
||||
return None
|
||||
|
||||
|
||||
def _safe_bool(value: object) -> bool | None:
|
||||
if value is None:
|
||||
return None
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue