런타임 계약과 학습자 흐름 보강
This commit is contained in:
parent
f456b8997a
commit
206018b088
56 changed files with 4306 additions and 1008 deletions
|
|
@ -14,6 +14,7 @@ from pydantic import BaseModel, Field
|
|||
from ..auth_types import AccountStatus, RoleName
|
||||
from ..auth_sessions import (
|
||||
ManagedUserPatch,
|
||||
ManagedUserUpsertInput,
|
||||
active_session_count,
|
||||
deactivate_managed_user,
|
||||
get_managed_user,
|
||||
|
|
@ -24,6 +25,7 @@ from ..auth_sessions import (
|
|||
upsert_managed_user,
|
||||
)
|
||||
from ..config import settings
|
||||
from ..contracts.engine_gateway import ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL
|
||||
from ..db import acquire, get_pool, healthcheck
|
||||
from ..deps import Principal, require_admin_access
|
||||
from ..engine_client import engine_client
|
||||
|
|
@ -955,7 +957,7 @@ def _default_engine_config() -> AdminEngineConfigResponse:
|
|||
return AdminEngineConfigResponse(
|
||||
engine_mode=settings.engine_mode,
|
||||
engine_url=settings.engine_url,
|
||||
model="gateway-default",
|
||||
model=ENGINE_GATEWAY_DEFAULT_MODEL_SENTINEL,
|
||||
durable=False,
|
||||
source="runtime_default",
|
||||
)
|
||||
|
|
@ -1725,14 +1727,16 @@ async def create_user(
|
|||
if body.role == "admin" or body.admin_access:
|
||||
_assert_super_admin(principal)
|
||||
user = await upsert_managed_user(
|
||||
email=_normalize_email(body.email),
|
||||
display_name=body.display_name,
|
||||
role=body.role,
|
||||
admin_access=body.admin_access,
|
||||
account_status=body.account_status,
|
||||
affiliation=body.affiliation,
|
||||
cohort_ids=body.cohort_ids,
|
||||
reactivate=True,
|
||||
ManagedUserUpsertInput(
|
||||
email=_normalize_email(body.email),
|
||||
display_name=body.display_name,
|
||||
role=body.role,
|
||||
admin_access=body.admin_access,
|
||||
account_status=body.account_status,
|
||||
affiliation=body.affiliation,
|
||||
cohort_ids=body.cohort_ids,
|
||||
reactivate=True,
|
||||
)
|
||||
)
|
||||
users, durable = await list_managed_users()
|
||||
if not durable:
|
||||
|
|
|
|||
|
|
@ -137,14 +137,11 @@ async def reevaluate_session(
|
|||
raise HTTPException(status.HTTP_503_SERVICE_UNAVAILABLE, detail=result.error)
|
||||
|
||||
await session_persistence.save_session_evaluation(
|
||||
session_id=session_id,
|
||||
learner_id=sess.learner_id,
|
||||
status="error" if result.error else "ready",
|
||||
source="engine",
|
||||
scope=result.scope,
|
||||
stage=result.stage,
|
||||
payload=result.to_dict(),
|
||||
error=result.error,
|
||||
session_persistence.SessionEvaluationWrite.from_result(
|
||||
session_id=session_id,
|
||||
learner_id=sess.learner_id,
|
||||
result=result,
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
|
|
@ -182,7 +179,7 @@ async def reevaluate_turn(
|
|||
client_reply = sess.turns[target_idx + 1].text_masked
|
||||
|
||||
# 평가용 경량 TurnContext 재구성(prepare_turn 의 결정론 산출과 동형). 엔진 호출 없음.
|
||||
from ..services.orchestrator import TurnContext # 지연 import(소유권 경계)
|
||||
from ..services.orchestrator import TurnContext, TurnMemory # 지연 import(소유권 경계)
|
||||
|
||||
recent = [
|
||||
{"speaker": tr.speaker, "text": tr.text_masked} for tr in sess.turns[max(0, target_idx - 4):target_idx]
|
||||
|
|
@ -195,7 +192,7 @@ async def reevaluate_turn(
|
|||
learner_text_raw=learner.text,
|
||||
learner_text_masked=learner.text_masked,
|
||||
state_after=sess.state, # 조회 시점 상태(정밀 재현은 DB 스냅샷 도입 시)
|
||||
recent_turns=recent,
|
||||
memory=TurnMemory(recent_turns=recent),
|
||||
)
|
||||
|
||||
result = await evaluator.evaluate_turn(
|
||||
|
|
|
|||
|
|
@ -13,6 +13,14 @@ from fastapi import APIRouter, Depends, HTTPException, Response, status
|
|||
from ..db import acquire
|
||||
from ..deps import CurrentPrincipal, Principal, Role, require_role
|
||||
from ..deps import AIView
|
||||
from ..persona_generation_contract import (
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT,
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE,
|
||||
coerce_persona_generated_draft,
|
||||
persona_draft_prompt_bundle,
|
||||
persona_generation_payload_from_response,
|
||||
persona_generation_schema,
|
||||
)
|
||||
from ..persona_repository import (
|
||||
archive_persona_family,
|
||||
create_persona_draft,
|
||||
|
|
@ -57,38 +65,6 @@ PERSONA_SOURCE_CITATION: dict[str, str] = {
|
|||
"textbook_guide": "교수자 첨부 교재/가이드 환언·발췌 근거 — 저작권 검수 필요",
|
||||
"mixed_notes": "교수자 첨부 혼합 메모 PII 마스킹 파생본",
|
||||
}
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_ID = "persona-draft-rag"
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_VERSION = "2026-06-28.1"
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT = (
|
||||
"출력은 반드시 structured_schema를 따른다. code는 P숫자 형식을 선호하되 "
|
||||
"힌트가 없으면 빈 문자열 대신 임시값 P로 둔다. source_provenance에는 "
|
||||
"RAG source_id와 첨부 근거 기반 초안임을 남긴다. evidence chunk id를 "
|
||||
"임상 필드 본문에 그대로 노출하지 않는다."
|
||||
)
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE = (
|
||||
"너는 Vignette 임상 콘텐츠 저작 보조자다. 아래 RAG 근거 청크만 바탕으로 교육용 "
|
||||
"가상내담자 페르소나 초안을 만든다. 첨부 원문은 KB 문서가 SSOT이며, 근거 밖 내용을 "
|
||||
"임의로 꾸며 핵심 임상 정보처럼 쓰지 않는다. 실제 개인정보는 이미 마스킹됐으며, "
|
||||
"원문 표현을 복사하지 말고 "
|
||||
"범주화·합성화된 임상 훈련용 설정으로 변환한다. CCD/DSM/역린은 런타임 내부 설정이므로 "
|
||||
"내담자 발화에 직접 노출되지 않는 형태로 작성한다."
|
||||
)
|
||||
|
||||
|
||||
def _persona_draft_prompt_bundle() -> dict[str, str]:
|
||||
payload = "\n".join(
|
||||
[
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_ID,
|
||||
PERSONA_DRAFT_PROMPT_BUNDLE_VERSION,
|
||||
PERSONA_DRAFT_SYSTEM_PROMPT,
|
||||
PERSONA_DRAFT_USER_PROMPT_PREAMBLE,
|
||||
]
|
||||
)
|
||||
return {
|
||||
"id": PERSONA_DRAFT_PROMPT_BUNDLE_ID,
|
||||
"version": PERSONA_DRAFT_PROMPT_BUNDLE_VERSION,
|
||||
"hash": hashlib.sha256(payload.encode("utf-8")).hexdigest()[:12],
|
||||
}
|
||||
|
||||
|
||||
def _card_from_draft_payload(request: PersonaDraftPayload):
|
||||
|
|
@ -454,142 +430,6 @@ def _format_generation_evidence(evidence: list[PersonaGenerationEvidence]) -> st
|
|||
return "\n\n".join(lines)
|
||||
|
||||
|
||||
def _persona_generation_schema() -> dict[str, Any]:
|
||||
return {
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"draft": {
|
||||
"type": "object",
|
||||
"additionalProperties": False,
|
||||
"properties": {
|
||||
"code": {"type": "string"},
|
||||
"display_name": {"type": "string"},
|
||||
"difficulty": {"type": "string", "enum": ["easy", "moderate", "hard"]},
|
||||
"theory_target": {"type": "array", "items": {"type": "string"}},
|
||||
"demographics": {"type": "object"},
|
||||
"presenting": {"type": "object"},
|
||||
"history": {"type": "object"},
|
||||
"big5": {"type": "object"},
|
||||
"resistance": {"type": "object"},
|
||||
"speech_style": {"type": "object"},
|
||||
"affect_baseline": {"type": "object"},
|
||||
"ccd": {"type": "object"},
|
||||
"dsm5_dimensional": {"type": "object"},
|
||||
"triggers": {"type": "object"},
|
||||
"source_provenance": {"type": "string"},
|
||||
"is_synthetic": {"type": "boolean"},
|
||||
},
|
||||
"required": [
|
||||
"code",
|
||||
"display_name",
|
||||
"difficulty",
|
||||
"theory_target",
|
||||
"demographics",
|
||||
"presenting",
|
||||
"history",
|
||||
"big5",
|
||||
"resistance",
|
||||
"speech_style",
|
||||
"affect_baseline",
|
||||
"ccd",
|
||||
"dsm5_dimensional",
|
||||
"triggers",
|
||||
"source_provenance",
|
||||
"is_synthetic",
|
||||
],
|
||||
},
|
||||
"source_summary": {"type": "string"},
|
||||
"warnings": {"type": "array", "items": {"type": "string"}},
|
||||
},
|
||||
"required": ["draft", "source_summary", "warnings"],
|
||||
}
|
||||
|
||||
|
||||
def _json_payload_from_generation(text: str) -> dict[str, Any]:
|
||||
try:
|
||||
parsed = json.loads(text)
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
except json.JSONDecodeError:
|
||||
start = text.find("{")
|
||||
end = text.rfind("}")
|
||||
if start >= 0 and end > start:
|
||||
try:
|
||||
parsed = json.loads(text[start : end + 1])
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
except json.JSONDecodeError:
|
||||
return {}
|
||||
return {}
|
||||
|
||||
|
||||
def _float_dict(value: Any) -> dict[str, float]:
|
||||
if not isinstance(value, dict):
|
||||
return {}
|
||||
result: dict[str, float] = {}
|
||||
for key, item in value.items():
|
||||
if isinstance(item, (int, float)):
|
||||
result[str(key)] = float(item)
|
||||
return result
|
||||
|
||||
|
||||
def _coerce_generated_draft(
|
||||
payload: dict[str, Any],
|
||||
request: PersonaDraftGenerateRequest,
|
||||
) -> PersonaDraftPayload:
|
||||
raw = payload.get("draft") if isinstance(payload.get("draft"), dict) else payload
|
||||
if not isinstance(raw, dict):
|
||||
raw = {}
|
||||
theory_target = raw.get("theory_target")
|
||||
theory_values = (
|
||||
[str(item).strip().lower() for item in theory_target if str(item).strip()]
|
||||
if isinstance(theory_target, list)
|
||||
else [value.strip().lower() for value in request.theory_target if value.strip()]
|
||||
)
|
||||
code = str(raw.get("code") or request.code_hint or "").strip().upper()
|
||||
display_name = str(raw.get("display_name") or request.display_name_hint or "자료 기반 새 페르소나").strip()
|
||||
difficulty = str(raw.get("difficulty") or request.difficulty)
|
||||
if difficulty not in {"easy", "moderate", "hard"}:
|
||||
difficulty = request.difficulty
|
||||
return PersonaDraftPayload(
|
||||
code=code or "P",
|
||||
display_name=display_name,
|
||||
difficulty=difficulty, # type: ignore[arg-type]
|
||||
theory_target=theory_values or ["humanistic"],
|
||||
demographics=_json_object(raw.get("demographics")),
|
||||
presenting=_json_object(raw.get("presenting")),
|
||||
history=_json_object(raw.get("history")),
|
||||
big5=_float_dict(raw.get("big5")) or {"O": 0.5, "C": 0.5, "E": 0.5, "A": 0.5, "N": 0.5},
|
||||
resistance=_float_dict(raw.get("resistance"))
|
||||
or {
|
||||
"base_resistance": 0.5,
|
||||
"unlock_rate": 0.1,
|
||||
"decay_floor": 0.05,
|
||||
"silence_prob": 0.15,
|
||||
"deflection_prob": 0.25,
|
||||
},
|
||||
speech_style=_json_object(raw.get("speech_style")),
|
||||
affect_baseline=_float_dict(raw.get("affect_baseline"))
|
||||
or {
|
||||
"negative_affect": 0.45,
|
||||
"hopelessness": 0.2,
|
||||
"anhedonia": 0.2,
|
||||
"sleep": 0.2,
|
||||
"anxiety": 0.35,
|
||||
"suicide_ideation_stage": 1,
|
||||
},
|
||||
ccd=_json_object(raw.get("ccd")),
|
||||
dsm5_dimensional=_json_object(raw.get("dsm5_dimensional")),
|
||||
triggers=_json_object(raw.get("triggers")),
|
||||
source_provenance=str(raw.get("source_provenance") or f"masked {request.source_kind}"),
|
||||
is_synthetic=bool(raw.get("is_synthetic", True)),
|
||||
submit_for_review=False,
|
||||
)
|
||||
|
||||
|
||||
def _json_object(value: Any) -> dict[str, Any]:
|
||||
return value if isinstance(value, dict) else {}
|
||||
|
||||
|
||||
def _ensure_teacher_or_admin(principal: Principal) -> None:
|
||||
if principal.role not in {Role.TEACHER, Role.ADMIN}:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail="only teachers and admins can review personas")
|
||||
|
|
@ -775,7 +615,7 @@ async def generate_persona_draft_route(
|
|||
query=evidence_query or "페르소나 저작 근거",
|
||||
)
|
||||
evidence_text = _format_generation_evidence(evidence)
|
||||
prompt_bundle = _persona_draft_prompt_bundle()
|
||||
prompt_bundle = persona_draft_prompt_bundle()
|
||||
prompt = (
|
||||
f"{PERSONA_DRAFT_USER_PROMPT_PREAMBLE}\n\n"
|
||||
f"자료 종류: {request.source_kind}\n"
|
||||
|
|
@ -799,7 +639,7 @@ async def generate_persona_draft_route(
|
|||
],
|
||||
max_tokens=2200,
|
||||
temperature=0.2,
|
||||
structured_schema=_persona_generation_schema(),
|
||||
structured_schema=persona_generation_schema(),
|
||||
metadata={
|
||||
"feature": "persona_draft_generation",
|
||||
"prompt_bundle": prompt_bundle,
|
||||
|
|
@ -814,8 +654,8 @@ async def generate_persona_draft_route(
|
|||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=f"persona draft generator unavailable: {exc}",
|
||||
) from exc
|
||||
payload = response.structured or _json_payload_from_generation(response.text)
|
||||
draft = _coerce_generated_draft(payload, request)
|
||||
payload = persona_generation_payload_from_response(response)
|
||||
draft = coerce_persona_generated_draft(payload, request)
|
||||
provenance = (
|
||||
f"prompt={prompt_bundle['id']}@{prompt_bundle['version']}#{prompt_bundle['hash']}; "
|
||||
f"RAG sources={','.join(source_ids)}; "
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime
|
||||
from typing import Literal, Optional
|
||||
|
|
@ -25,7 +26,16 @@ from ..deps import CurrentPrincipal, Principal, Role
|
|||
from ..engine_client import EngineError, engine_client
|
||||
from ..persona_repository import get_catalog_persona
|
||||
from ..runtime_policy import require_runtime_fallback_allowed
|
||||
from ..services import evaluator, guardrail, live_coach, memory, orchestrator, rag, state_machine
|
||||
from ..services import (
|
||||
evaluator,
|
||||
guardrail,
|
||||
live_coach,
|
||||
memory,
|
||||
orchestrator,
|
||||
rag,
|
||||
session_digest_worker,
|
||||
state_machine,
|
||||
)
|
||||
from ..session_read_model import (
|
||||
LearnerDashboardResponse,
|
||||
LearnerSessionsResponse,
|
||||
|
|
@ -58,6 +68,7 @@ from ..session_read_model import (
|
|||
from ..store import InProcSession, TurnRecord, store
|
||||
|
||||
router = APIRouter(prefix="/sessions", tags=["sessions"])
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TheoryMode = Literal["humanistic", "cbt", "integrative"]
|
||||
EndStateValue = str | int | float | bool | None | dict[str, float]
|
||||
|
|
@ -394,10 +405,12 @@ async def _prepare_turn_context(
|
|||
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
|
||||
|
|
@ -502,12 +515,52 @@ async def _end_persisted_session(sess: InProcSession, carry: memory.CarryOver) -
|
|||
sess.ended = True
|
||||
sess.ended_at = datetime.now().timestamp()
|
||||
store.put(sess)
|
||||
if _should_schedule_session_digest_worker(carry):
|
||||
asyncio.create_task(_run_session_digest_worker_for_session(sess.session_id))
|
||||
asyncio.create_task(_write_episodic_embeddings(sess))
|
||||
return
|
||||
require_runtime_fallback_allowed("session end")
|
||||
store.end(sess.session_id)
|
||||
|
||||
|
||||
def _should_schedule_session_digest_worker(carry: memory.CarryOver) -> bool:
|
||||
return bool(settings.session_digest_worker_enabled and carry.compression_job is not None)
|
||||
|
||||
|
||||
async def _run_session_digest_worker_for_session(session_id: str) -> None:
|
||||
"""Best-effort M2 LLM digest compressor.
|
||||
|
||||
The DB connection is held only for load/apply. Engine generation runs outside
|
||||
the transaction so a slow provider cannot pin the pool.
|
||||
"""
|
||||
|
||||
try:
|
||||
db.get_pool()
|
||||
async with db.acquire(role="admin") as conn:
|
||||
loaded = await session_digest_worker.load_session_digest_job(conn, session_id)
|
||||
if loaded is None:
|
||||
return
|
||||
model = settings.session_digest_worker_model.strip() or None
|
||||
worker = await session_digest_worker.run_session_digest_worker(
|
||||
loaded.job,
|
||||
engine_client,
|
||||
existing_case_digest=loaded.existing_case_digest,
|
||||
model=model,
|
||||
audit_hook=session_persistence.record_llm_call_audit,
|
||||
)
|
||||
if worker.apply_plan is None:
|
||||
return
|
||||
async with db.acquire(role="admin") as conn:
|
||||
await session_digest_worker.apply_session_digest_plan(
|
||||
conn,
|
||||
worker.apply_plan,
|
||||
learner_id=loaded.learner_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.warning("session digest worker failed for session_id=%s", session_id, exc_info=True)
|
||||
return
|
||||
|
||||
|
||||
async def _write_episodic_embeddings(sess: InProcSession) -> None:
|
||||
"""Best-effort M2 episodic writer.
|
||||
|
||||
|
|
@ -601,27 +654,22 @@ async def _generate_and_save_session_evaluation(sess: InProcSession) -> None:
|
|||
),
|
||||
timeout=min(float(settings.engine_timeout), 45.0),
|
||||
)
|
||||
status_value = "error" if result.error else "ready"
|
||||
await session_persistence.save_session_evaluation(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
status=status_value,
|
||||
source="engine",
|
||||
scope=result.scope,
|
||||
stage=result.stage,
|
||||
payload=result.to_dict(),
|
||||
error=result.error,
|
||||
session_persistence.SessionEvaluationWrite.from_result(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
result=result,
|
||||
)
|
||||
)
|
||||
except Exception as exc:
|
||||
await session_persistence.save_session_evaluation(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
status="error",
|
||||
source="engine",
|
||||
scope="session_end",
|
||||
stage=_stage_label(sess.state.stage),
|
||||
payload={},
|
||||
error=str(exc),
|
||||
session_persistence.SessionEvaluationWrite.from_error(
|
||||
session_id=sess.session_id,
|
||||
learner_id=sess.learner_id,
|
||||
scope="session_end",
|
||||
stage=_stage_label(sess.state.stage),
|
||||
error=str(exc),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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