대시보드 폴드아웃/드릴다운 정리 + 페르소나 역린·misconduct 반응 + 게이트웨이 격리·RAG 비차단 수정
SSOT 대시보드:
- 한신대 기술분석 PDF(19쪽) 정합성 분석 + 이번 세션 발견 섹션 추가
- 섹션 폴드아웃(접기)·상단 목차(드릴다운)·모두 펼치기/접기 — 내용 보존, 레이아웃만 정리
페르소나 반응 강화('저항·반응 조절' 핵심 차별):
- PersonaCard.triggers(역린) 필드 + CCD 핵심상처 파생 역린 블록
- L0에 무례·모욕·조롱 시 현실적 동맹 균열 반응 지침
버그·성능 수정(라이브/E2E로 포착):
- 게이트웨이 페르소나 격리: --append-system-prompt를 --system-prompt(교체)로 + --exclude-dynamic-system-prompt-sections (내담자 캐릭터 붕괴·개발맥락 누출 차단)
- RAG: 임베더 동기 로드(약 7-13초)를 _warm_rag_caches 백그라운드 warm으로(세션 생성 블로킹 회귀 수정)
- voice TTS RMS 데드힌트 제거, init_state OpennessParams 파라미터객체화
- 한국어 PII(날짜·금액·주소) 마스킹 보강
- 레이아웃 시각 게이트: 폼 컨트롤 값 스크롤 오탐 제외(7/7)
검증: 백엔드 84/84, E2E 42(데스크톱 27·모바일 11·아바타 4), 시각 게이트 7/7
This commit is contained in:
parent
cb2aebd76c
commit
085460b5e0
327 changed files with 31226 additions and 1829 deletions
|
|
@ -27,6 +27,19 @@ def _is_local_url(value: str) -> bool:
|
|||
return host in {"localhost", "127.0.0.1", "::1"}
|
||||
|
||||
|
||||
def _is_allowed_local_dev_cors_origin(value: str) -> bool:
|
||||
parsed = urlsplit(value)
|
||||
host = (parsed.hostname or "").lower()
|
||||
return (
|
||||
parsed.scheme == "http"
|
||||
and host in {"localhost", "127.0.0.1"}
|
||||
and parsed.port in range(5170, 5181)
|
||||
and not parsed.path
|
||||
and not parsed.query
|
||||
and not parsed.fragment
|
||||
)
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
|
|
@ -106,6 +119,22 @@ class Settings(BaseSettings):
|
|||
default=False,
|
||||
validation_alias="AUTH_DEV_LOGIN_ENABLED",
|
||||
)
|
||||
auth_saml_enabled: bool = Field(
|
||||
default=False,
|
||||
validation_alias="AUTH_SAML_ENABLED",
|
||||
)
|
||||
saml_sp_entity_id: str = Field(
|
||||
default="",
|
||||
validation_alias="SAML_SP_ENTITY_ID",
|
||||
)
|
||||
saml_sso_url: str = Field(
|
||||
default="",
|
||||
validation_alias="SAML_SSO_URL",
|
||||
)
|
||||
saml_x509_cert_fingerprint: str = Field(
|
||||
default="",
|
||||
validation_alias="SAML_X509_CERT_FINGERPRINT",
|
||||
)
|
||||
default_affiliation: str = Field(
|
||||
default="",
|
||||
validation_alias="DEFAULT_AFFILIATION",
|
||||
|
|
@ -160,11 +189,23 @@ class Settings(BaseSettings):
|
|||
forbidden.append("SESSION_SECRET")
|
||||
if _is_local_url(self.frontend_base_url):
|
||||
forbidden.append("FRONTEND_BASE_URL")
|
||||
if any(_is_local_url(origin) for origin in self.cors_origins):
|
||||
if any(
|
||||
_is_local_url(origin) and not _is_allowed_local_dev_cors_origin(origin)
|
||||
for origin in self.cors_origins
|
||||
):
|
||||
forbidden.append("CORS_ORIGINS")
|
||||
if forbidden:
|
||||
joined = ", ".join(forbidden)
|
||||
raise ValueError(f"{joined} must be production-safe when ENVIRONMENT={self.environment}")
|
||||
if self.auth_saml_enabled:
|
||||
missing_saml: list[str] = []
|
||||
if not self.saml_sp_entity_id.strip():
|
||||
missing_saml.append("SAML_SP_ENTITY_ID")
|
||||
if not self.saml_sso_url.strip():
|
||||
missing_saml.append("SAML_SSO_URL")
|
||||
if missing_saml:
|
||||
joined = ", ".join(missing_saml)
|
||||
raise ValueError(f"{joined} must be configured when AUTH_SAML_ENABLED=true")
|
||||
return self
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -39,8 +39,6 @@ class GenerateRequest(BaseModel):
|
|||
|
||||
ai_role: AIRole
|
||||
messages: list[EngineMessage]
|
||||
# tier 라우팅 힌트: client=Sonnet/Solar, evaluator=Opus, fast=Haiku (마스터플랜 §5)
|
||||
tier: Literal["client", "feedback", "fast"] = "client"
|
||||
model: Optional[str] = None # 명시 시 게이트웨이 override
|
||||
max_tokens: int = 1024
|
||||
temperature: float = 0.7
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from __future__ import annotations
|
|||
import json
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Iterable
|
||||
from typing import Any, Iterable, Literal, cast
|
||||
|
||||
from .config import settings
|
||||
from .db import acquire, get_pool
|
||||
|
|
@ -18,6 +18,7 @@ from .services.persona import PersonaCard, SEED_PERSONAS, get_seed_persona
|
|||
|
||||
|
||||
SEED_VERSION = 1
|
||||
PersonaStatus = Literal["draft", "review", "approved", "archived"]
|
||||
|
||||
_CARD_COLUMNS = """
|
||||
persona_id, code, version, status, display_name, difficulty, theory_target,
|
||||
|
|
@ -25,6 +26,15 @@ _CARD_COLUMNS = """
|
|||
affect_baseline, ccd, dsm5_dimensional, source_provenance, is_synthetic
|
||||
"""
|
||||
|
||||
_REVIEW_COLUMNS = """
|
||||
persona_id, code, version, status, display_name, difficulty, theory_target,
|
||||
source_provenance, is_synthetic, created_at, approved_at
|
||||
"""
|
||||
|
||||
_PERSONA_STATUSES = {"draft", "review", "approved", "archived"}
|
||||
_REVIEW_QUEUE_STATUSES = ("draft", "review")
|
||||
PersonaReviewAction = Literal["approve", "reject"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CatalogPersona:
|
||||
|
|
@ -35,6 +45,21 @@ class CatalogPersona:
|
|||
degraded: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PersonaReviewItem:
|
||||
persona_id: str
|
||||
code: str
|
||||
version: int
|
||||
status: PersonaStatus
|
||||
display_name: str
|
||||
difficulty: str
|
||||
theory_target: list[str]
|
||||
source_provenance: str
|
||||
is_synthetic: bool
|
||||
created_at: str | None
|
||||
approved_at: str | None
|
||||
|
||||
|
||||
def seed_persona_id(code: str) -> str:
|
||||
return str(uuid.uuid5(uuid.NAMESPACE_URL, f"vignette:persona:{code.upper()}"))
|
||||
|
||||
|
|
@ -54,6 +79,24 @@ def _string_list(value: Iterable[Any] | None) -> list[str]:
|
|||
return [str(item) for item in value]
|
||||
|
||||
|
||||
def _optional_text(value: Any) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
isoformat = getattr(value, "isoformat", None)
|
||||
if callable(isoformat):
|
||||
return str(isoformat())
|
||||
return str(value)
|
||||
|
||||
|
||||
def _normalize_statuses(statuses: Iterable[str]) -> list[str]:
|
||||
normalized: list[str] = []
|
||||
for status in statuses:
|
||||
value = str(status).strip().lower()
|
||||
if value in _PERSONA_STATUSES and value not in normalized:
|
||||
normalized.append(value)
|
||||
return normalized
|
||||
|
||||
|
||||
def card_from_row(row: Any) -> PersonaCard:
|
||||
return PersonaCard(
|
||||
code=str(row["code"]).upper(),
|
||||
|
|
@ -84,6 +127,22 @@ def catalog_persona_from_row(row: Any) -> CatalogPersona:
|
|||
)
|
||||
|
||||
|
||||
def persona_review_item_from_row(row: Any) -> PersonaReviewItem:
|
||||
return PersonaReviewItem(
|
||||
persona_id=str(row["persona_id"]),
|
||||
code=str(row["code"]).upper(),
|
||||
version=int(row["version"]),
|
||||
status=cast(PersonaStatus, str(row["status"]).lower()),
|
||||
display_name=str(row["display_name"]),
|
||||
difficulty=str(row["difficulty"]),
|
||||
theory_target=_string_list(row["theory_target"]),
|
||||
source_provenance=str(row["source_provenance"] or ""),
|
||||
is_synthetic=bool(row["is_synthetic"]),
|
||||
created_at=_optional_text(row["created_at"]),
|
||||
approved_at=_optional_text(row["approved_at"]),
|
||||
)
|
||||
|
||||
|
||||
def seed_fallback_persona(code: str) -> CatalogPersona | None:
|
||||
card = get_seed_persona(code)
|
||||
if card is None:
|
||||
|
|
@ -209,6 +268,93 @@ async def get_approved_persona(code: str) -> CatalogPersona | None:
|
|||
return catalog_persona_from_row(row) if row is not None else None
|
||||
|
||||
|
||||
async def list_persona_review_queue(
|
||||
*,
|
||||
role: str,
|
||||
statuses: Iterable[str] = _REVIEW_QUEUE_STATUSES,
|
||||
) -> list[PersonaReviewItem]:
|
||||
if role not in {"teacher", "admin"}:
|
||||
raise ValueError("persona review queue requires teacher or admin role")
|
||||
status_values = _normalize_statuses(statuses)
|
||||
if not status_values:
|
||||
return []
|
||||
|
||||
get_pool()
|
||||
async with acquire(role=role) as conn:
|
||||
rows = await conn.fetch(
|
||||
f"""
|
||||
SELECT {_REVIEW_COLUMNS}
|
||||
FROM app.persona_card
|
||||
WHERE status = ANY($1::text[])
|
||||
ORDER BY
|
||||
CASE status
|
||||
WHEN 'review' THEN 0
|
||||
WHEN 'draft' THEN 1
|
||||
ELSE 2
|
||||
END,
|
||||
code,
|
||||
version DESC
|
||||
""",
|
||||
status_values,
|
||||
)
|
||||
return [persona_review_item_from_row(row) for row in rows]
|
||||
|
||||
|
||||
async def update_persona_review_status(
|
||||
*,
|
||||
persona_id: str,
|
||||
action: PersonaReviewAction,
|
||||
reviewer_id: str,
|
||||
role: str,
|
||||
) -> PersonaReviewItem | None:
|
||||
if role not in {"teacher", "admin"}:
|
||||
raise ValueError("persona review update requires teacher or admin role")
|
||||
if action not in {"approve", "reject"}:
|
||||
raise ValueError("unsupported persona review action")
|
||||
|
||||
next_status = "approved" if action == "approve" else "draft"
|
||||
approved_by = reviewer_id if action == "approve" else None
|
||||
approved_at_expr = "now()" if action == "approve" else "NULL"
|
||||
|
||||
get_pool()
|
||||
async with acquire(role=role, user_id=reviewer_id) as conn:
|
||||
row = await conn.fetchrow(
|
||||
f"""
|
||||
UPDATE app.persona_card
|
||||
SET
|
||||
status = $2,
|
||||
approved_by = $3::uuid,
|
||||
approved_at = {approved_at_expr}
|
||||
WHERE persona_id = $1::uuid
|
||||
AND status IN ('draft', 'review')
|
||||
RETURNING {_REVIEW_COLUMNS}
|
||||
""",
|
||||
persona_id,
|
||||
next_status,
|
||||
approved_by,
|
||||
)
|
||||
if row is None:
|
||||
return None
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO audit.audit_log (
|
||||
actor_uid, action, target_kind, target_id, detail
|
||||
)
|
||||
VALUES ($1::uuid, $2, $3, $4, $5::jsonb)
|
||||
""",
|
||||
reviewer_id,
|
||||
f"persona_{action}",
|
||||
"persona_card",
|
||||
persona_id,
|
||||
{
|
||||
"next_status": next_status,
|
||||
"code": str(row["code"]).upper(),
|
||||
"version": int(row["version"]),
|
||||
},
|
||||
)
|
||||
return persona_review_item_from_row(row)
|
||||
|
||||
|
||||
async def list_catalog_personas() -> list[CatalogPersona]:
|
||||
try:
|
||||
return await list_approved_personas()
|
||||
|
|
@ -229,6 +375,9 @@ async def get_catalog_persona(code: str) -> CatalogPersona | None:
|
|||
|
||||
__all__ = [
|
||||
"CatalogPersona",
|
||||
"PersonaReviewItem",
|
||||
"PersonaReviewAction",
|
||||
"PersonaStatus",
|
||||
"SEED_VERSION",
|
||||
"card_from_row",
|
||||
"catalog_persona_from_row",
|
||||
|
|
@ -236,8 +385,11 @@ __all__ = [
|
|||
"get_catalog_persona",
|
||||
"list_approved_personas",
|
||||
"list_catalog_personas",
|
||||
"list_persona_review_queue",
|
||||
"materialize_seed_personas",
|
||||
"persona_review_item_from_row",
|
||||
"seed_fallback_persona",
|
||||
"seed_fallback_personas",
|
||||
"seed_persona_id",
|
||||
"update_persona_review_status",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -25,6 +25,13 @@ from pydantic import BaseModel
|
|||
from ..auth_sessions import InactiveUserError, SessionUser, create_session, revoke_session
|
||||
from ..config import settings
|
||||
from ..deps import CurrentPrincipal, Principal, Role
|
||||
from ..saml import (
|
||||
SamlIdentity,
|
||||
acs_url_for_entity_id,
|
||||
build_authn_request,
|
||||
parse_fixture_response,
|
||||
redirect_binding_url,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/auth", tags=["auth"])
|
||||
|
||||
|
|
@ -41,7 +48,15 @@ class OAuthState:
|
|||
created_at: float
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SamlState:
|
||||
request_id: str
|
||||
next_path: str
|
||||
created_at: float
|
||||
|
||||
|
||||
_oauth_states: dict[str, OAuthState] = {}
|
||||
_saml_states: dict[str, SamlState] = {}
|
||||
|
||||
|
||||
class MeResponse(BaseModel):
|
||||
|
|
@ -52,8 +67,17 @@ class MeResponse(BaseModel):
|
|||
cohort_ids: list[str]
|
||||
|
||||
|
||||
class AuthProviderStatus(BaseModel):
|
||||
provider: Literal["google", "saml"]
|
||||
configured: bool
|
||||
enabled: bool
|
||||
login_path: str
|
||||
|
||||
|
||||
class AuthConfigResponse(BaseModel):
|
||||
google_oauth_configured: bool
|
||||
saml_configured: bool
|
||||
providers: list[AuthProviderStatus]
|
||||
allowed_email_domains: list[str]
|
||||
redirect_uri: str
|
||||
dev_login_enabled: bool
|
||||
|
|
@ -84,6 +108,37 @@ def _normalize_email_set(values: list[str]) -> set[str]:
|
|||
return {email for value in values if (email := _normalize_email(value))}
|
||||
|
||||
|
||||
def _google_configured() -> bool:
|
||||
return bool(settings.oauth_google_client_id and settings.oauth_google_client_secret)
|
||||
|
||||
|
||||
def _saml_configured() -> bool:
|
||||
return bool(
|
||||
settings.auth_saml_enabled
|
||||
and settings.saml_sp_entity_id.strip()
|
||||
and settings.saml_sso_url.strip()
|
||||
)
|
||||
|
||||
|
||||
def _auth_provider_statuses() -> list[AuthProviderStatus]:
|
||||
google_ready = _google_configured()
|
||||
saml_ready = _saml_configured()
|
||||
return [
|
||||
AuthProviderStatus(
|
||||
provider="google",
|
||||
configured=google_ready,
|
||||
enabled=google_ready,
|
||||
login_path="/auth/login?provider=google",
|
||||
),
|
||||
AuthProviderStatus(
|
||||
provider="saml",
|
||||
configured=saml_ready,
|
||||
enabled=saml_ready,
|
||||
login_path="/auth/login?provider=saml",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def allowed_email_domains() -> set[str]:
|
||||
"""Configured login email domains, normalized for claim checks."""
|
||||
return {
|
||||
|
|
@ -137,6 +192,15 @@ def _role_for_email(email: str) -> Role:
|
|||
return Role.LEARNER
|
||||
|
||||
|
||||
def _role_for_saml_identity(identity: SamlIdentity) -> Role:
|
||||
hinted = (identity.role_hint or "").strip().lower()
|
||||
if hinted in {"admin", "administrator"}:
|
||||
return Role.ADMIN
|
||||
if hinted in {"teacher", "instructor", "faculty"}:
|
||||
return Role.TEACHER
|
||||
return _role_for_email(identity.email)
|
||||
|
||||
|
||||
def _safe_next_path(next_path: str | None) -> str:
|
||||
if not next_path or not next_path.startswith("/") or next_path.startswith("//"):
|
||||
return "/"
|
||||
|
|
@ -211,6 +275,13 @@ def _prune_oauth_states() -> None:
|
|||
_oauth_states.pop(key, None)
|
||||
|
||||
|
||||
def _prune_saml_states() -> None:
|
||||
cutoff = time.time() - OAUTH_STATE_TTL_SECONDS
|
||||
stale = [key for key, value in _saml_states.items() if value.created_at < cutoff]
|
||||
for key in stale:
|
||||
_saml_states.pop(key, None)
|
||||
|
||||
|
||||
def _cookie_secure() -> bool:
|
||||
# The __Host- prefix requires Secure, Path=/, and no Domain. Modern Chrome
|
||||
# accepts Secure cookies on localhost, which keeps dev and prod semantics
|
||||
|
|
@ -291,10 +362,12 @@ def _dev_login_available(request: Request) -> bool:
|
|||
@router.get("/config", response_model=AuthConfigResponse)
|
||||
async def auth_config(request: Request) -> AuthConfigResponse:
|
||||
"""Return non-secret login configuration for the browser login screen."""
|
||||
google_ready = _google_configured()
|
||||
saml_ready = _saml_configured()
|
||||
return AuthConfigResponse(
|
||||
google_oauth_configured=bool(
|
||||
settings.oauth_google_client_id and settings.oauth_google_client_secret
|
||||
),
|
||||
google_oauth_configured=google_ready,
|
||||
saml_configured=saml_ready,
|
||||
providers=_auth_provider_statuses(),
|
||||
allowed_email_domains=sorted(allowed_email_domains()),
|
||||
redirect_uri=settings.oauth_redirect_uri,
|
||||
dev_login_enabled=_dev_login_available(request),
|
||||
|
|
@ -308,9 +381,33 @@ async def login(
|
|||
next: Annotated[str | None, Query()] = None,
|
||||
) -> RedirectResponse:
|
||||
"""Start Google OIDC authorization code + PKCE login."""
|
||||
if provider == "saml":
|
||||
if not _saml_configured():
|
||||
return _frontend_login_redirect("saml_not_configured", request)
|
||||
_prune_saml_states()
|
||||
relay_state = secrets.token_urlsafe(32)
|
||||
acs_url = acs_url_for_entity_id(settings.saml_sp_entity_id)
|
||||
request_id, authn_request_xml = build_authn_request(
|
||||
sp_entity_id=settings.saml_sp_entity_id,
|
||||
sso_url=settings.saml_sso_url,
|
||||
acs_url=acs_url,
|
||||
)
|
||||
_saml_states[relay_state] = SamlState(
|
||||
request_id=request_id,
|
||||
next_path=_safe_next_path(next),
|
||||
created_at=time.time(),
|
||||
)
|
||||
return RedirectResponse(
|
||||
redirect_binding_url(
|
||||
sso_url=settings.saml_sso_url,
|
||||
authn_request_xml=authn_request_xml,
|
||||
relay_state=relay_state,
|
||||
),
|
||||
status_code=302,
|
||||
)
|
||||
if provider != "google":
|
||||
return _frontend_login_redirect("unsupported_provider", request)
|
||||
if not settings.oauth_google_client_id or not settings.oauth_google_client_secret:
|
||||
if not _google_configured():
|
||||
return _frontend_login_redirect("not_configured", request)
|
||||
|
||||
_prune_oauth_states()
|
||||
|
|
@ -406,6 +503,58 @@ async def callback(
|
|||
return response
|
||||
|
||||
|
||||
@router.post("/saml/acs")
|
||||
async def saml_acs(request: Request) -> RedirectResponse:
|
||||
"""Accept a minimal unsigned SAMLResponse for local fixture SAML proof.
|
||||
|
||||
Signed SAML verification is intentionally not implemented. When
|
||||
SAML_X509_CERT_FINGERPRINT is configured, this endpoint refuses to trust the
|
||||
response so production does not silently run unsigned SAML.
|
||||
"""
|
||||
if not _saml_configured():
|
||||
return _frontend_login_redirect("saml_not_configured", request)
|
||||
if settings.saml_x509_cert_fingerprint.strip():
|
||||
return _frontend_login_redirect("saml_signature_verification_required", request)
|
||||
if settings.environment != "dev":
|
||||
return _frontend_login_redirect("saml_fixture_acs_dev_only", request)
|
||||
|
||||
form = await request.form()
|
||||
relay_state = str(form.get("RelayState") or "")
|
||||
encoded_response = str(form.get("SAMLResponse") or "")
|
||||
if not relay_state or not encoded_response:
|
||||
return _frontend_login_redirect("saml_missing_callback", request)
|
||||
|
||||
_prune_saml_states()
|
||||
stored = _saml_states.pop(relay_state, None)
|
||||
if stored is None:
|
||||
return _frontend_login_redirect("saml_invalid_state", request)
|
||||
|
||||
try:
|
||||
identity = parse_fixture_response(encoded_response)
|
||||
email = validate_google_identity_domain(
|
||||
email=identity.email,
|
||||
email_verified=True,
|
||||
hosted_domain=_email_domain(identity.email),
|
||||
)
|
||||
except (HTTPException, ValueError):
|
||||
return _frontend_login_redirect("saml_assertion_invalid", request)
|
||||
|
||||
role = _role_for_saml_identity(identity)
|
||||
try:
|
||||
sid, _ = await create_session(
|
||||
email=email,
|
||||
display_name=identity.display_name or email,
|
||||
role=role.value,
|
||||
cohort_ids=[],
|
||||
)
|
||||
except InactiveUserError:
|
||||
return _frontend_login_redirect("inactive_user", request)
|
||||
|
||||
response = RedirectResponse(_frontend_url(stored.next_path, request), status_code=302)
|
||||
_set_session_cookie(response, sid)
|
||||
return response
|
||||
|
||||
|
||||
@router.post("/dev-login", response_model=MeResponse)
|
||||
async def dev_login(request: Request, body: DevLoginRequest, response: Response) -> MeResponse:
|
||||
"""Dev-only server login for local E2E and manual testing.
|
||||
|
|
|
|||
|
|
@ -2,15 +2,23 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from typing import Annotated, Any, Literal
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Response, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Response, status
|
||||
from pydantic import BaseModel
|
||||
|
||||
from ..deps import CurrentPrincipal
|
||||
from ..persona_repository import CatalogPersona, list_catalog_personas
|
||||
from ..deps import CurrentPrincipal, Principal, Role, require_role
|
||||
from ..persona_repository import (
|
||||
CatalogPersona,
|
||||
PersonaReviewAction,
|
||||
PersonaReviewItem,
|
||||
list_catalog_personas,
|
||||
list_persona_review_queue,
|
||||
update_persona_review_status,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/personas", tags=["personas"])
|
||||
TeacherOrAdmin = Annotated[Principal, Depends(require_role(Role.TEACHER, Role.ADMIN))]
|
||||
|
||||
|
||||
class PersonaSummary(BaseModel):
|
||||
|
|
@ -25,6 +33,24 @@ class PersonaSummary(BaseModel):
|
|||
degraded: bool = False
|
||||
|
||||
|
||||
class PersonaReviewSummary(BaseModel):
|
||||
persona_id: str
|
||||
code: str
|
||||
version: int
|
||||
status: Literal["draft", "review", "approved", "archived"]
|
||||
display_name: str
|
||||
difficulty: str
|
||||
theory_target: list[str]
|
||||
source_provenance: str
|
||||
is_synthetic: bool
|
||||
created_at: str | None = None
|
||||
approved_at: str | None = None
|
||||
|
||||
|
||||
class PersonaReviewDecisionRequest(BaseModel):
|
||||
action: PersonaReviewAction
|
||||
|
||||
|
||||
def _first_text_value(data: dict[str, Any]) -> str:
|
||||
for value in data.values():
|
||||
if isinstance(value, str) and value.strip():
|
||||
|
|
@ -46,6 +72,27 @@ def _summary(entry: CatalogPersona) -> PersonaSummary:
|
|||
)
|
||||
|
||||
|
||||
def _review_summary(entry: PersonaReviewItem) -> PersonaReviewSummary:
|
||||
return PersonaReviewSummary(
|
||||
persona_id=entry.persona_id,
|
||||
code=entry.code,
|
||||
version=entry.version,
|
||||
status=entry.status,
|
||||
display_name=entry.display_name,
|
||||
difficulty=entry.difficulty,
|
||||
theory_target=entry.theory_target,
|
||||
source_provenance=entry.source_provenance,
|
||||
is_synthetic=entry.is_synthetic,
|
||||
created_at=entry.created_at,
|
||||
approved_at=entry.approved_at,
|
||||
)
|
||||
|
||||
|
||||
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")
|
||||
|
||||
|
||||
@router.get("", response_model=list[PersonaSummary])
|
||||
async def list_personas(response: Response, _principal: CurrentPrincipal) -> list[PersonaSummary]:
|
||||
"""Return latest approved personas from app.persona_card."""
|
||||
|
|
@ -64,3 +111,47 @@ async def list_personas(response: Response, _principal: CurrentPrincipal) -> lis
|
|||
response.headers["X-Vignette-Catalog-Source"] = "database"
|
||||
|
||||
return [_summary(entry) for entry in personas]
|
||||
|
||||
|
||||
@router.get("/review", response_model=list[PersonaReviewSummary])
|
||||
async def list_persona_reviews(principal: TeacherOrAdmin) -> list[PersonaReviewSummary]:
|
||||
"""Return draft/review personas awaiting faculty approval."""
|
||||
_ensure_teacher_or_admin(principal)
|
||||
try:
|
||||
queue = await list_persona_review_queue(role=principal.role.value)
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="persona review queue database unavailable",
|
||||
) from exc
|
||||
return [_review_summary(entry) for entry in queue]
|
||||
|
||||
|
||||
@router.post("/review/{persona_id}", response_model=PersonaReviewSummary)
|
||||
async def decide_persona_review(
|
||||
persona_id: str,
|
||||
request: PersonaReviewDecisionRequest,
|
||||
principal: TeacherOrAdmin,
|
||||
) -> PersonaReviewSummary:
|
||||
"""Approve a persona for learners or return it to draft for changes."""
|
||||
_ensure_teacher_or_admin(principal)
|
||||
try:
|
||||
updated = await update_persona_review_status(
|
||||
persona_id=persona_id,
|
||||
action=request.action,
|
||||
reviewer_id=principal.user_id,
|
||||
role=principal.role.value,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status.HTTP_403_FORBIDDEN, detail=str(exc)) from exc
|
||||
except Exception as exc:
|
||||
raise HTTPException(
|
||||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail="persona review update database unavailable",
|
||||
) from exc
|
||||
if updated is None:
|
||||
raise HTTPException(
|
||||
status.HTTP_404_NOT_FOUND,
|
||||
detail="persona review item not found or not pending review",
|
||||
)
|
||||
return _review_summary(updated)
|
||||
|
|
|
|||
|
|
@ -18,13 +18,13 @@ from fastapi import APIRouter, HTTPException, status
|
|||
from pydantic import BaseModel, Field
|
||||
from sse_starlette.sse import EventSourceResponse
|
||||
|
||||
from .. import session_persistence
|
||||
from .. import db, session_persistence
|
||||
from ..config import settings
|
||||
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, runtime_fallback_allowed
|
||||
from ..services import evaluator, memory, orchestrator, state_machine
|
||||
from ..services import evaluator, memory, orchestrator, rag, state_machine
|
||||
from ..store import InProcSession, TurnRecord, store
|
||||
|
||||
router = APIRouter(prefix="/sessions", tags=["sessions"])
|
||||
|
|
@ -193,6 +193,165 @@ class SessionReviewResponse(BaseModel):
|
|||
|
||||
|
||||
_RECALL_CACHE: dict[str, memory.RecallContext] = {}
|
||||
# 세션별 KB 증상 행동단서(회기 1회 산출·캐시). 빈 list 캐시 = 회기 내 재시도 안 함(안정성).
|
||||
_KB_CUES_CACHE: dict[str, list[str]] = {}
|
||||
_LEARNER_VISIBLE_AI_ROLE = "counselor"
|
||||
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
# RAG 배선 헬퍼 — 내담자(CLIENT) 뷰. 임베더/KB/DB 풀 미가용 시 빈 값으로 graceful
|
||||
# degradation: 상담 루프를 절대 막지 않는다(라이브 루프 비차단이 계약). routes/kb.py가
|
||||
# 같은 예외를 503으로 올리는 것과 의도적으로 다르다. 임베딩은 rag가 스레드풀로 offload.
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
_RAG_RECALL_K = 5
|
||||
_KB_CUES_K = 4
|
||||
|
||||
|
||||
def _persona_kb_query(card) -> str:
|
||||
"""페르소나 증상·호소 → KB 행동단서 검색 질의(임베더/tsquery 입력 전용, LLM 미주입).
|
||||
|
||||
질의는 프롬프트에 들어가지 않는다. 회수된 behavior_cue만 L2로 주입되고, CLIENT 정책
|
||||
(expose_body=False)이 본문을 잘라 '행동단서'만 돌려준다(CCD 본문 비노출 자동 보존).
|
||||
"""
|
||||
parts: list[str] = []
|
||||
presenting = getattr(card, "presenting", None) or {}
|
||||
if presenting.get("주호소"):
|
||||
parts.append(str(presenting["주호소"]))
|
||||
if presenting.get("표층"):
|
||||
parts.append(str(presenting["표층"]))
|
||||
dsm = getattr(card, "dsm5_dimensional", None) or {}
|
||||
parts.extend(str(key) for key in dsm.keys() if key != "note")
|
||||
return " ".join(p for p in parts if p).strip()
|
||||
|
||||
|
||||
async def _retrieve_kb_behavior_cues(card) -> list[str]:
|
||||
"""KB 증상 행동단서 회수(CLIENT 정책). 미가용 시 빈 리스트(비차단)."""
|
||||
query = _persona_kb_query(card)
|
||||
if not query:
|
||||
return []
|
||||
try:
|
||||
async with db.acquire(ai_view=rag.AIRole.CLIENT.value) as conn:
|
||||
result = await rag.search_kb(
|
||||
conn,
|
||||
query=query,
|
||||
role=rag.AIRole.CLIENT,
|
||||
k=_KB_CUES_K,
|
||||
)
|
||||
return [c.behavior_cue for c in result.chunks if c.behavior_cue]
|
||||
except Exception:
|
||||
# rag.NotConfigured(임베더/KB 미가용)·RuntimeError(풀 미초기화)·DB 오류 포함.
|
||||
# 비치명적: 빈 단서로 진행. CancelledError는 BaseException이라 미포착.
|
||||
return []
|
||||
|
||||
|
||||
async def _ensure_kb_cues(session_id: str, card) -> list[str]:
|
||||
"""세션별 KB 행동단서(회기 1회 산출·캐시, 서버 재시작/재개 시 lazy 재계산)."""
|
||||
cached = _KB_CUES_CACHE.get(session_id)
|
||||
if cached is not None:
|
||||
return cached
|
||||
cues = await _retrieve_kb_behavior_cues(card)
|
||||
_KB_CUES_CACHE[session_id] = cues
|
||||
return cues
|
||||
|
||||
|
||||
async def _load_prev_case_summary(case_id: str) -> Optional[dict]:
|
||||
"""직전 회기 요약(case 스코프) → build_recall_context 입력. 미존재/미가용 시 None."""
|
||||
try:
|
||||
async with db.acquire(ai_view=rag.AIRole.CLIENT.value) as conn:
|
||||
row = await conn.fetchrow(
|
||||
"""
|
||||
SELECT digest, open_threads, end_state
|
||||
FROM app.session_summary
|
||||
WHERE case_id = $1::uuid
|
||||
ORDER BY session_no DESC, created_at DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
case_id,
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
if row is None:
|
||||
return None
|
||||
return {
|
||||
"digest": row["digest"],
|
||||
"open_threads": list(row["open_threads"] or []),
|
||||
"end_state": dict(row["end_state"] or {}),
|
||||
}
|
||||
|
||||
|
||||
async def _hydrate_episodic_text(conn, result) -> list[str]:
|
||||
"""retrieve_persona_memory가 돌려준 turn_id → app.turns 마스킹 본문 조인(내담자 발화)."""
|
||||
turn_ids = [c.meta.get("turn_id") for c in result.chunks if c.meta.get("turn_id")]
|
||||
if not turn_ids:
|
||||
return []
|
||||
rows = await conn.fetch(
|
||||
"""
|
||||
SELECT id, text_masked FROM app.turns
|
||||
WHERE id = ANY($1::uuid[]) AND speaker = 'client'
|
||||
""",
|
||||
turn_ids,
|
||||
)
|
||||
by_id = {str(r["id"]): r["text_masked"] for r in rows}
|
||||
return [by_id[t] for t in turn_ids if by_id.get(t)]
|
||||
|
||||
|
||||
def _recall_query(prev_summary: Optional[dict], card) -> str:
|
||||
"""episodic recall 질의: 직전 open_threads 우선, 없으면 주호소."""
|
||||
if prev_summary:
|
||||
threads = prev_summary.get("open_threads") or []
|
||||
if threads:
|
||||
return " ".join(str(t) for t in threads)
|
||||
presenting = getattr(card, "presenting", None) or {}
|
||||
return str(presenting.get("주호소") or "").strip()
|
||||
|
||||
|
||||
async def _episodic_recall_snippets(case_id: str, query: str) -> list[str]:
|
||||
"""case 스코프 episodic 벡터 recall → 내담자 발화 단편(마스킹본). 미가용 시 []."""
|
||||
if not query:
|
||||
return []
|
||||
try:
|
||||
async with db.acquire(ai_view=rag.AIRole.CLIENT.value) as conn:
|
||||
result = await rag.retrieve_persona_memory(
|
||||
conn, case_id=case_id, query=query, k=_RAG_RECALL_K,
|
||||
)
|
||||
return await _hydrate_episodic_text(conn, result)
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
|
||||
async def _build_start_recall(*, case_id: str, card) -> memory.RecallContext:
|
||||
"""회기 시작 회상 조립: prev_summary(case) + episodic recall을 build_recall_context로
|
||||
합본. 전 구간 graceful(미가용 시 빈 회상).
|
||||
"""
|
||||
try:
|
||||
db.get_pool() # 풀 미초기화 시 RuntimeError → 첫 회기와 동일한 빈 회상
|
||||
except RuntimeError:
|
||||
return memory.build_recall_context()
|
||||
prev_summary = await _load_prev_case_summary(case_id)
|
||||
query = _recall_query(prev_summary, card)
|
||||
episodic = await _episodic_recall_snippets(case_id, query)
|
||||
pinned = list((prev_summary or {}).get("pinned_facts") or [])
|
||||
return memory.build_recall_context(
|
||||
prev_summary=prev_summary,
|
||||
episodic_snippets=episodic,
|
||||
pinned_facts=pinned,
|
||||
)
|
||||
|
||||
|
||||
async def _warm_rag_caches(session_id: str, case_id: str, card) -> None:
|
||||
"""RAG 회상·KB 행동단서를 **백그라운드**로 산출해 캐시한다(요청 경로 비차단).
|
||||
|
||||
BGE-M3 임베더 첫 로드(~수 초)가 회기 시작/턴 응답을 막지 않도록 create_task로 띄운다.
|
||||
warm 완료 전 턴은 빈 회상/단서로 진행(graceful), 이후 턴부터 RAG 주입. 전 구간 비치명적.
|
||||
"""
|
||||
try:
|
||||
_RECALL_CACHE[session_id] = await _build_start_recall(case_id=case_id, card=card)
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
_KB_CUES_CACHE[session_id] = await _retrieve_kb_behavior_cues(card)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
_PHASE_KEY_BY_LABEL = {
|
||||
"라포": "rapport",
|
||||
|
|
@ -481,6 +640,92 @@ def _evaluation_payload(record: dict[str, object] | None) -> dict[str, object]:
|
|||
return payload if isinstance(payload, dict) else {}
|
||||
|
||||
|
||||
# fast-loop 턴 평가(TechniqueCategory) → 프론트 sr-technique--{kind} 시각 매핑.
|
||||
_TECHNIQUE_KIND_BY_CATEGORY = {
|
||||
"relational": "empathy",
|
||||
"exploratory": "explore",
|
||||
"intervention": "confront",
|
||||
"stabilizing": "reflect",
|
||||
"structuring": "closed",
|
||||
}
|
||||
|
||||
|
||||
def _review_techniques_from_turn_eval(ev: dict[str, object] | None) -> list[ReviewTechnique]:
|
||||
"""턴 평가의 기법 태그를 리뷰 칩으로. label_ko 우선, category로 색 kind 결정."""
|
||||
if not isinstance(ev, dict):
|
||||
return []
|
||||
out: list[ReviewTechnique] = []
|
||||
for tag in ev.get("techniques") or []:
|
||||
if not isinstance(tag, dict):
|
||||
continue
|
||||
label = str(tag.get("label_ko") or tag.get("code") or "").strip()
|
||||
if not label:
|
||||
continue
|
||||
kind = _TECHNIQUE_KIND_BY_CATEGORY.get(str(tag.get("category") or ""), "explore")
|
||||
out.append(ReviewTechnique(kind=kind, label=label))
|
||||
return out
|
||||
|
||||
|
||||
def _review_note_from_turn_eval(ev: dict[str, object] | None) -> Optional[ReviewNote]:
|
||||
"""의도이탈(있으면 우선) 또는 적절성 신호를 턴 노트로. tone: good|warn(프론트 계약)."""
|
||||
if not isinstance(ev, dict):
|
||||
return None
|
||||
dev = ev.get("intent_deviation")
|
||||
if isinstance(dev, dict):
|
||||
dimension = str(dev.get("dimension") or "").strip()
|
||||
expected = str(dev.get("expected") or "").strip()
|
||||
actual = str(dev.get("actual") or "").strip()
|
||||
body = " / ".join(p for p in (f"권장: {expected}" if expected else "", f"실제: {actual}" if actual else "") if p)
|
||||
return ReviewNote(
|
||||
author="평가 AI",
|
||||
tone="warn",
|
||||
title=f"의도와 다른 부분 · {dimension}".rstrip(" ·") or "의도와 다른 부분",
|
||||
body=body or "권장 반응과 실제 반응에 차이가 있었어요.",
|
||||
)
|
||||
appropriateness = str(ev.get("appropriateness") or "neutral")
|
||||
note_text = str(ev.get("appropriateness_note") or "").strip()
|
||||
if appropriateness == "pos":
|
||||
return ReviewNote(author="평가 AI", tone="good", title="적절한 개입", body=note_text or "이 개입은 흐름에 적절했어요.")
|
||||
if appropriateness == "warn" and note_text:
|
||||
return ReviewNote(author="평가 AI", tone="warn", title="점검해볼 지점", body=note_text)
|
||||
return None
|
||||
|
||||
|
||||
async def _record_safety_event(sess: InProcSession, ctx, result) -> None:
|
||||
"""위기 escalate 시 app.safety_events 적재(교수자 감사·알림 레코드). C2.
|
||||
|
||||
비차단: DB 미가용(degraded)·FK 미충족(in-memory 세션) 시 graceful skip — 상담 루프를
|
||||
절대 막지 않는다. 실시간 교수자 push 알림은 후속(이 레코드가 1차 알림원).
|
||||
"""
|
||||
crisis = getattr(ctx, "crisis", None)
|
||||
if crisis is None or not getattr(crisis, "escalate", False):
|
||||
return
|
||||
kind = getattr(crisis.kind, "value", None) or str(getattr(crisis, "kind", "crisis"))
|
||||
try:
|
||||
async with db.acquire() as conn:
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO app.safety_events
|
||||
(session_id, trigger_type, ko_risk_level, escalated, detail)
|
||||
VALUES ($1::uuid, $2, $3, TRUE, $4::jsonb)
|
||||
""",
|
||||
sess.session_id,
|
||||
kind,
|
||||
int(getattr(crisis, "risk_level", 0) or 0),
|
||||
json.dumps({
|
||||
"matched": list(getattr(crisis, "matched", []) or []),
|
||||
"stage": getattr(result, "stage", None),
|
||||
"turn_seq": getattr(result, "turn_seq", None),
|
||||
}),
|
||||
)
|
||||
except Exception:
|
||||
pass # 비차단(R5): 적재 실패가 위기 대응/상담을 막지 않음.
|
||||
|
||||
|
||||
def _learner_visible_turns(sess: InProcSession) -> list[TurnRecord]:
|
||||
return sess.turns_visible_to(_LEARNER_VISIBLE_AI_ROLE)
|
||||
|
||||
|
||||
async def _generate_and_save_session_evaluation(sess: InProcSession) -> None:
|
||||
if not sess.turns:
|
||||
return
|
||||
|
|
@ -535,8 +780,9 @@ def _schedule_session_evaluation(sess: InProcSession) -> None:
|
|||
|
||||
|
||||
def _learner_summary(sess: InProcSession, *, review_ready: bool = False) -> LearnerSessionSummary:
|
||||
learner_turns = sum(1 for turn in sess.turns if turn.speaker == "counselor")
|
||||
client_turns = sum(1 for turn in sess.turns if turn.speaker == "client")
|
||||
turns = _learner_visible_turns(sess)
|
||||
learner_turns = sum(1 for turn in turns if turn.speaker == "counselor")
|
||||
client_turns = sum(1 for turn in turns if turn.speaker == "client")
|
||||
return LearnerSessionSummary(
|
||||
session_id=sess.session_id,
|
||||
persona_code=sess.persona_code,
|
||||
|
|
@ -544,7 +790,7 @@ def _learner_summary(sess: InProcSession, *, review_ready: bool = False) -> Lear
|
|||
session_no=sess.session_no,
|
||||
status="ended" if sess.ended else "active",
|
||||
stage=_stage_label(sess.state.stage),
|
||||
turn_count=len(sess.turns),
|
||||
turn_count=len(turns),
|
||||
learner_turn_count=learner_turns,
|
||||
client_turn_count=client_turns,
|
||||
started_at=_iso(sess.created_at) or "",
|
||||
|
|
@ -554,7 +800,10 @@ def _learner_summary(sess: InProcSession, *, review_ready: bool = False) -> Lear
|
|||
|
||||
|
||||
async def _review_ready(sess: InProcSession, principal: Principal) -> bool:
|
||||
if not sess.ended or not sess.turns:
|
||||
turns = _learner_visible_turns(sess)
|
||||
if not sess.ended or not turns:
|
||||
return False
|
||||
if len(turns) != len(sess.turns):
|
||||
return False
|
||||
evaluation_record, _ = await session_persistence.load_session_evaluation(
|
||||
sess.session_id,
|
||||
|
|
@ -568,6 +817,7 @@ def _session_detail(
|
|||
*,
|
||||
review_ready: bool = False,
|
||||
) -> SessionDetailResponse:
|
||||
turns = _learner_visible_turns(sess)
|
||||
return SessionDetailResponse(
|
||||
session_id=sess.session_id,
|
||||
case_id=sess.case_id,
|
||||
|
|
@ -587,7 +837,7 @@ def _session_detail(
|
|||
text=turn.text_masked,
|
||||
created_at=_iso(turn.created_at) or "",
|
||||
)
|
||||
for turn in sess.turns
|
||||
for turn in turns
|
||||
],
|
||||
review_ready=review_ready,
|
||||
)
|
||||
|
|
@ -649,10 +899,7 @@ async def start_session(
|
|||
|
||||
recall = memory.build_recall_context()
|
||||
st = state_machine.init_state(
|
||||
base_resistance=card.base_resistance(),
|
||||
unlock_rate=card.unlock_rate(),
|
||||
decay_floor=card.decay_floor(),
|
||||
ideation_baseline=card.ideation_baseline(),
|
||||
params=card.openness_params(),
|
||||
carry=recall.carry,
|
||||
)
|
||||
|
||||
|
|
@ -680,7 +927,11 @@ async def start_session(
|
|||
)
|
||||
else:
|
||||
store.put(sess)
|
||||
|
||||
# 즉시 빈/carry 회상으로 응답을 막지 않는다. RAG 회상·KB 단서(임베더 로드 수 초)는
|
||||
# 백그라운드 warm으로 캐시 — 회기 시작/턴 응답이 임베더 로드에 블로킹되지 않게(성능 회귀 방지).
|
||||
_RECALL_CACHE[sess.session_id] = recall
|
||||
asyncio.create_task(_warm_rag_caches(sess.session_id, sess.case_id, card))
|
||||
|
||||
return SessionStartResponse(
|
||||
session_id=sess.session_id,
|
||||
|
|
@ -701,6 +952,8 @@ async def get_session_review(
|
|||
"""Return a learner-safe review built only from the stored session transcript."""
|
||||
_ensure_learner(principal)
|
||||
sess = await _load_session_or_404(session_id, principal, allow_ended=True)
|
||||
visible_turns = _learner_visible_turns(sess)
|
||||
hidden_turns = len(visible_turns) != len(sess.turns)
|
||||
|
||||
end_ts = sess.ended_at or datetime.now().timestamp()
|
||||
duration_seconds = max(0, int(round(end_ts - sess.created_at)))
|
||||
|
|
@ -708,7 +961,7 @@ async def get_session_review(
|
|||
client_initial = client_name[:1] or "내"
|
||||
|
||||
reached_phase = _stage_label(sess.state.stage)
|
||||
stage_labels = [turn.stage for turn in sess.turns] or [reached_phase]
|
||||
stage_labels = [turn.stage for turn in visible_turns] or [reached_phase]
|
||||
axis = ["0:00"]
|
||||
if duration_seconds > 0:
|
||||
axis.append(_offset_label(duration_seconds))
|
||||
|
|
@ -717,16 +970,20 @@ async def get_session_review(
|
|||
session_id,
|
||||
principal,
|
||||
)
|
||||
evaluation_payload = _evaluation_payload(evaluation_record)
|
||||
evaluation_status = str(evaluation_record.get("status") or "") if evaluation_record else ""
|
||||
evaluation_ready = evaluation_status == "ready"
|
||||
evaluation_payload = {} if hidden_turns else _evaluation_payload(evaluation_record)
|
||||
evaluation_status = (
|
||||
"" if hidden_turns else str(evaluation_record.get("status") or "") if evaluation_record else ""
|
||||
)
|
||||
evaluation_ready = not hidden_turns and evaluation_status == "ready"
|
||||
|
||||
first_turn_ts = sess.turns[0].created_at if sess.turns else sess.created_at
|
||||
first_turn_ts = visible_turns[0].created_at if visible_turns else sess.created_at
|
||||
turns: list[ReviewTurn] = []
|
||||
for index, turn in enumerate(sess.turns):
|
||||
for index, turn in enumerate(visible_turns):
|
||||
speaker: Literal["learner", "client"] = (
|
||||
"learner" if turn.speaker == "counselor" else "client"
|
||||
)
|
||||
# 턴별 fast-loop 평가는 학습자 발화에만 부착(기법 태깅·노트). hidden 시 노출 안 함.
|
||||
turn_eval = turn.evaluation if (speaker == "learner" and not hidden_turns) else None
|
||||
turns.append(
|
||||
ReviewTurn(
|
||||
id=f"t{index + 1}",
|
||||
|
|
@ -734,8 +991,8 @@ async def get_session_review(
|
|||
speaker=speaker,
|
||||
who="학습자" if speaker == "learner" else client_name,
|
||||
text=turn.text_masked,
|
||||
techniques=[],
|
||||
note=None,
|
||||
techniques=_review_techniques_from_turn_eval(turn_eval),
|
||||
note=_review_note_from_turn_eval(turn_eval),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -783,10 +1040,10 @@ async def get_session_review(
|
|||
|
||||
summary = _review_summary_from_evaluation(
|
||||
fallback=transcript_summary,
|
||||
evaluation_record=evaluation_record,
|
||||
evaluation_record=None if hidden_turns else evaluation_record,
|
||||
payload=evaluation_payload,
|
||||
)
|
||||
if evaluation_record and not evaluation_durable:
|
||||
if evaluation_record and not hidden_turns and not evaluation_durable:
|
||||
summary += " 현재 평가는 런타임 캐시에서 복원되었습니다."
|
||||
|
||||
return SessionReviewResponse(
|
||||
|
|
@ -832,6 +1089,7 @@ async def submit_turn(
|
|||
_ensure_learner(principal)
|
||||
sess = await _load_session_or_404(session_id, principal)
|
||||
recall = _RECALL_CACHE.get(session_id) or memory.RecallContext()
|
||||
kb_cues = _KB_CUES_CACHE.get(session_id) or [] # 비차단: warm 전이면 빈 단서(graceful)
|
||||
|
||||
ctx = orchestrator.prepare_turn(
|
||||
session_id=session_id,
|
||||
|
|
@ -841,18 +1099,25 @@ async def submit_turn(
|
|||
learner_text=body.text,
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(),
|
||||
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
|
||||
|
||||
try:
|
||||
result = await orchestrator.run_turn_generate(ctx, engine_client)
|
||||
result = await orchestrator.run_turn_generate(
|
||||
ctx,
|
||||
engine_client,
|
||||
eval_hook=evaluator.make_eval_hook(engine_client),
|
||||
)
|
||||
except EngineError as exc:
|
||||
raise HTTPException(
|
||||
status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=f"engine unavailable: {exc}",
|
||||
) from exc
|
||||
|
||||
# 턴별 fast-loop 평가는 학습자(상담자) 발화에 부착(기법 태깅·적절성·의도이탈).
|
||||
await _append_session_turn(
|
||||
sess,
|
||||
TurnRecord(
|
||||
|
|
@ -861,6 +1126,7 @@ async def submit_turn(
|
|||
stage=_stage_label(ctx.state_after.stage),
|
||||
text=body.text,
|
||||
text_masked=ctx.learner_text_masked,
|
||||
evaluation=result.evaluation,
|
||||
),
|
||||
)
|
||||
|
||||
|
|
@ -873,9 +1139,15 @@ async def submit_turn(
|
|||
stage=_stage_label(result.state_after.stage),
|
||||
text=result.client_reply,
|
||||
text_masked=result.client_reply,
|
||||
llm_provider=result.llm_provider,
|
||||
model=result.model,
|
||||
tokens_in=result.tokens_in,
|
||||
tokens_out=result.tokens_out,
|
||||
cost_usd=result.cost_usd,
|
||||
),
|
||||
)
|
||||
await _update_session_state(sess, result.state_after)
|
||||
await _record_safety_event(sess, ctx, result) # C2: 위기 escalate 시 safety_events 적재(비차단)
|
||||
|
||||
return TurnResponse(
|
||||
turn_seq=result.turn_seq,
|
||||
|
|
@ -897,6 +1169,7 @@ async def stream_turn(
|
|||
_ensure_learner(principal)
|
||||
sess = await _load_session_or_404(session_id, principal)
|
||||
recall = _RECALL_CACHE.get(session_id) or memory.RecallContext()
|
||||
kb_cues = _KB_CUES_CACHE.get(session_id) or [] # 비차단: warm 전이면 빈 단서(graceful)
|
||||
|
||||
ctx = orchestrator.prepare_turn(
|
||||
session_id=session_id,
|
||||
|
|
@ -906,7 +1179,9 @@ async def stream_turn(
|
|||
learner_text=body.text,
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(),
|
||||
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
|
||||
|
||||
|
|
@ -941,6 +1216,11 @@ async def stream_turn(
|
|||
stage=_stage_label(ctx.state_after.stage),
|
||||
text=final_reply,
|
||||
text_masked=final_reply,
|
||||
llm_provider=str(ev.data.get("llm_provider") or ""),
|
||||
model=str(ev.data.get("model") or ""),
|
||||
tokens_in=int(ev.data.get("tokens_in") or 0),
|
||||
tokens_out=int(ev.data.get("tokens_out") or 0),
|
||||
cost_usd=float(ev.data.get("cost_usd") or 0.0),
|
||||
),
|
||||
)
|
||||
yield {"event": "done", "data": json.dumps(data, ensure_ascii=False)}
|
||||
|
|
@ -980,6 +1260,7 @@ async def end_session(
|
|||
|
||||
await _end_persisted_session(sess, carry)
|
||||
_RECALL_CACHE.pop(session_id, None)
|
||||
_KB_CUES_CACHE.pop(session_id, None)
|
||||
_schedule_session_evaluation(sess)
|
||||
|
||||
return SessionEndResponse(
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ cleanly instead of crashing.
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import hashlib
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
|
||||
|
|
@ -27,7 +29,7 @@ from ..deps import Principal, Role
|
|||
from ..engine_client import EngineError, engine_client
|
||||
from ..persona_repository import get_catalog_persona
|
||||
from ..runtime_policy import require_runtime_fallback_allowed, runtime_fallback_allowed
|
||||
from ..services import memory, orchestrator, state_machine
|
||||
from ..services import evaluator, memory, orchestrator, state_machine
|
||||
from ..services import voice as voice_svc
|
||||
from ..services.voice import VoicePreset, VoiceUnavailable, resolve_voice, voice_service
|
||||
from ..store import InProcSession, TurnRecord, store
|
||||
|
|
@ -113,6 +115,8 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
|
||||
audio_buf = bytearray()
|
||||
receiving = False
|
||||
audio_started_at: float | None = None
|
||||
last_audio_end_at: float | None = None
|
||||
|
||||
try:
|
||||
while True:
|
||||
|
|
@ -126,6 +130,7 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
if not receiving:
|
||||
# Be tolerant when audio arrives before audio_start.
|
||||
receiving = True
|
||||
audio_started_at = time.monotonic()
|
||||
audio_buf.clear()
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "listening"})
|
||||
audio_buf.extend(msg["bytes"])
|
||||
|
|
@ -151,11 +156,16 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
ctype = ctrl.get("type")
|
||||
if ctype == "audio_start":
|
||||
receiving = True
|
||||
audio_started_at = time.monotonic()
|
||||
audio_buf.clear()
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "listening"})
|
||||
|
||||
elif ctype == "audio_end":
|
||||
receiving = False
|
||||
audio_ended_at = time.monotonic()
|
||||
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))
|
||||
await _handle_utterance(
|
||||
websocket,
|
||||
session_id=session_id,
|
||||
|
|
@ -163,7 +173,13 @@ async def voice_ws(websocket: WebSocket) -> None:
|
|||
voice_preset=voice_preset,
|
||||
audio=bytes(audio_buf),
|
||||
fmt=ctrl.get("format"),
|
||||
audio_started_at=audio_started_at,
|
||||
audio_ended_at=audio_ended_at,
|
||||
silence_ms=silence_ms,
|
||||
barge_in=_safe_bool(ctrl.get("barge_in")),
|
||||
)
|
||||
last_audio_end_at = audio_ended_at
|
||||
audio_started_at = None
|
||||
audio_buf.clear()
|
||||
|
||||
elif ctype == "text_turn":
|
||||
|
|
@ -202,6 +218,10 @@ async def _handle_utterance(
|
|||
voice_preset: VoicePreset,
|
||||
audio: bytes,
|
||||
fmt: Optional[str],
|
||||
audio_started_at: float | None = None,
|
||||
audio_ended_at: float | None = None,
|
||||
silence_ms: int | None = None,
|
||||
barge_in: bool | None = None,
|
||||
) -> None:
|
||||
"""Transcribe one utterance, generate the client reply, then synthesize TTS."""
|
||||
if not audio:
|
||||
|
|
@ -226,6 +246,9 @@ async def _handle_utterance(
|
|||
return
|
||||
|
||||
learner_text = stt.text
|
||||
audio_ref = _voice_audio_ref(audio, fmt)
|
||||
duration_s = stt.duration or _elapsed_seconds(audio_started_at, audio_ended_at)
|
||||
speech_rate = _estimate_speech_rate(learner_text, duration_s)
|
||||
await _safe_send_json(
|
||||
websocket,
|
||||
{"type": "transcript", "text": learner_text, "final": True, "speaker": "counselor"},
|
||||
|
|
@ -240,6 +263,10 @@ async def _handle_utterance(
|
|||
principal=principal,
|
||||
voice_preset=voice_preset,
|
||||
learner_text=learner_text,
|
||||
audio_ref=audio_ref,
|
||||
silence_ms=silence_ms,
|
||||
speech_rate=speech_rate,
|
||||
barge_in=barge_in,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -250,6 +277,10 @@ async def _run_turn_and_speak(
|
|||
principal: Principal,
|
||||
voice_preset: VoicePreset,
|
||||
learner_text: str,
|
||||
audio_ref: str | None = None,
|
||||
silence_ms: int | None = None,
|
||||
speech_rate: float | None = None,
|
||||
barge_in: bool | None = None,
|
||||
) -> None:
|
||||
"""Run one counseling turn and stream synthesized client speech."""
|
||||
sess, err = await _load_voice_session(session_id, principal)
|
||||
|
|
@ -267,13 +298,18 @@ async def _run_turn_and_speak(
|
|||
learner_text=learner_text,
|
||||
recall_summary=recall.recall_summary,
|
||||
pinned_facts=recall.pinned_facts,
|
||||
recent_turns=sess.recent_turns(),
|
||||
recent_turns=sess.recent_turns(visible_to="client"),
|
||||
theory_mode=sess.theory_mode,
|
||||
)
|
||||
assert ctx.state_after is not None
|
||||
|
||||
# Voice needs the full client reply before TTS starts.
|
||||
try:
|
||||
result = await orchestrator.run_turn_generate(ctx, engine_client)
|
||||
result = await orchestrator.run_turn_generate(
|
||||
ctx,
|
||||
engine_client,
|
||||
eval_hook=evaluator.make_eval_hook(engine_client),
|
||||
)
|
||||
except EngineError as e:
|
||||
await _safe_send_json(websocket, {"type": "error", "detail": f"engine unavailable: {e}"})
|
||||
await _safe_send_json(websocket, {"type": "state", "state": "idle"})
|
||||
|
|
@ -290,6 +326,11 @@ async def _run_turn_and_speak(
|
|||
stage=ctx.state_after.stage.value,
|
||||
text=learner_text,
|
||||
text_masked=ctx.learner_text_masked,
|
||||
audio_ref=audio_ref,
|
||||
silence_ms=silence_ms,
|
||||
speech_rate=speech_rate,
|
||||
barge_in=barge_in,
|
||||
evaluation=result.evaluation,
|
||||
),
|
||||
)
|
||||
if reply:
|
||||
|
|
@ -302,6 +343,11 @@ async def _run_turn_and_speak(
|
|||
stage=result.stage,
|
||||
text=reply,
|
||||
text_masked=reply,
|
||||
llm_provider=result.llm_provider,
|
||||
model=result.model,
|
||||
tokens_in=result.tokens_in,
|
||||
tokens_out=result.tokens_out,
|
||||
cost_usd=result.cost_usd,
|
||||
),
|
||||
)
|
||||
await _update_voice_state(sess, result.state_after)
|
||||
|
|
@ -333,10 +379,7 @@ async def _run_turn_and_speak(
|
|||
try:
|
||||
n = 0
|
||||
async for ck in voice_service.synthesize_stream(reply, voice_preset):
|
||||
# Metadata precedes the binary chunk so the client can pair them.
|
||||
await _safe_send_json(
|
||||
websocket, {"type": "tts_chunk", "seq": ck.seq, "rms": round(ck.rms, 4)}
|
||||
)
|
||||
# 바이너리 오디오 청크만 송신(프론트가 Web Audio AnalyserNode로 립싱크 자체 산출).
|
||||
await _safe_send_bytes(websocket, ck.audio)
|
||||
n += 1
|
||||
await _safe_send_json(websocket, {"type": "tts_end", "chunks": n})
|
||||
|
|
@ -451,10 +494,7 @@ async def _bind_session(
|
|||
card = catalog_persona.card
|
||||
|
||||
st = state_machine.init_state(
|
||||
base_resistance=card.base_resistance(),
|
||||
unlock_rate=card.unlock_rate(),
|
||||
decay_floor=card.decay_floor(),
|
||||
ideation_baseline=card.ideation_baseline(),
|
||||
params=card.openness_params(),
|
||||
)
|
||||
sess = await session_persistence.create_session(
|
||||
learner_id=principal.user_id,
|
||||
|
|
@ -511,6 +551,52 @@ def _audio_meta(fmt: Optional[str]) -> tuple[str, str]:
|
|||
return table.get(f, ("audio.webm", "audio/webm"))
|
||||
|
||||
|
||||
def _voice_audio_ref(audio: bytes, fmt: Optional[str]) -> str | None:
|
||||
if not audio:
|
||||
return None
|
||||
f = (fmt or "webm").lower().lstrip(".") or "webm"
|
||||
digest = hashlib.sha256(audio).hexdigest()[:24]
|
||||
return f"voice:{f}:sha256:{digest}"
|
||||
|
||||
|
||||
def _elapsed_seconds(started_at: float | None, ended_at: float | None) -> float | None:
|
||||
if started_at is None or ended_at is None:
|
||||
return None
|
||||
return max(0.001, ended_at - started_at)
|
||||
|
||||
|
||||
def _estimate_speech_rate(text: str, duration_s: float | None) -> float | None:
|
||||
if not text or not duration_s or duration_s <= 0:
|
||||
return None
|
||||
units = sum(1 for ch in text if not ch.isspace())
|
||||
if units <= 0:
|
||||
return None
|
||||
return round((units / duration_s) * 60.0, 2)
|
||||
|
||||
|
||||
def _safe_int(value: object) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _safe_bool(value: object) -> bool | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in {"1", "true", "yes", "y"}:
|
||||
return True
|
||||
if normalized in {"0", "false", "no", "n"}:
|
||||
return False
|
||||
return bool(value)
|
||||
|
||||
|
||||
async def _safe_send_json(websocket: WebSocket, payload: dict) -> None:
|
||||
if websocket.client_state != WebSocketState.CONNECTED:
|
||||
return
|
||||
|
|
|
|||
157
apps/api/app/saml.py
Normal file
157
apps/api/app/saml.py
Normal file
|
|
@ -0,0 +1,157 @@
|
|||
"""Minimal SAML SP helpers for local fixture authentication tests.
|
||||
|
||||
This module intentionally implements only the Redirect-binding AuthnRequest and
|
||||
unsigned fixture ACS parsing needed for backend proof. Signed production SAML
|
||||
assertion verification is not implemented here.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import html
|
||||
import uuid
|
||||
import zlib
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Iterable
|
||||
from urllib.parse import urlsplit, urlunsplit, urlencode
|
||||
from xml.etree import ElementTree
|
||||
|
||||
|
||||
SAML_PROTOCOL_NS = "urn:oasis:names:tc:SAML:2.0:protocol"
|
||||
SAML_ASSERTION_NS = "urn:oasis:names:tc:SAML:2.0:assertion"
|
||||
SAML_ATTRIBUTE_ROLE_NAMES = {
|
||||
"role",
|
||||
"roles",
|
||||
"groups",
|
||||
"memberOf",
|
||||
"http://schemas.microsoft.com/ws/2008/06/identity/claims/role",
|
||||
}
|
||||
SAML_ATTRIBUTE_EMAIL_NAMES = {
|
||||
"email",
|
||||
"mail",
|
||||
"emailaddress",
|
||||
"EmailAddress",
|
||||
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/emailaddress",
|
||||
}
|
||||
SAML_ATTRIBUTE_DISPLAY_NAME_NAMES = {
|
||||
"display_name",
|
||||
"displayName",
|
||||
"name",
|
||||
"cn",
|
||||
"http://schemas.xmlsoap.org/ws/2005/05/identity/claims/name",
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SamlIdentity:
|
||||
email: str
|
||||
display_name: str
|
||||
role_hint: str | None = None
|
||||
|
||||
|
||||
def acs_url_for_entity_id(entity_id: str) -> str:
|
||||
parsed = urlsplit(entity_id.strip())
|
||||
if parsed.scheme and parsed.netloc:
|
||||
path = parsed.path.rstrip("/")
|
||||
if path.endswith("/metadata"):
|
||||
path = path[: -len("/metadata")]
|
||||
return urlunsplit((parsed.scheme, parsed.netloc, f"{path}/acs", "", ""))
|
||||
value = entity_id.strip().rstrip("/")
|
||||
if value.endswith("/metadata"):
|
||||
value = value[: -len("/metadata")]
|
||||
return value + "/acs"
|
||||
|
||||
|
||||
def build_authn_request(
|
||||
*,
|
||||
sp_entity_id: str,
|
||||
sso_url: str,
|
||||
acs_url: str,
|
||||
) -> tuple[str, str]:
|
||||
request_id = "_" + uuid.uuid4().hex
|
||||
issued_at = datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
xml = (
|
||||
f'<samlp:AuthnRequest xmlns:samlp="{SAML_PROTOCOL_NS}" '
|
||||
f'xmlns:saml="{SAML_ASSERTION_NS}" ID="{request_id}" Version="2.0" '
|
||||
f'IssueInstant="{issued_at}" Destination="{html.escape(sso_url, quote=True)}" '
|
||||
f'AssertionConsumerServiceURL="{html.escape(acs_url, quote=True)}" '
|
||||
f'ProtocolBinding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST">'
|
||||
f"<saml:Issuer>{html.escape(sp_entity_id)}</saml:Issuer>"
|
||||
"</samlp:AuthnRequest>"
|
||||
)
|
||||
return request_id, xml
|
||||
|
||||
|
||||
def redirect_binding_url(*, sso_url: str, authn_request_xml: str, relay_state: str) -> str:
|
||||
compressor = zlib.compressobj(wbits=-15)
|
||||
deflated = compressor.compress(authn_request_xml.encode("utf-8")) + compressor.flush()
|
||||
params = urlencode(
|
||||
{
|
||||
"SAMLRequest": base64.b64encode(deflated).decode("ascii"),
|
||||
"RelayState": relay_state,
|
||||
}
|
||||
)
|
||||
separator = "&" if "?" in sso_url else "?"
|
||||
return f"{sso_url}{separator}{params}"
|
||||
|
||||
|
||||
def inflate_redirect_request(encoded_request: str) -> str:
|
||||
payload = base64.b64decode(encoded_request)
|
||||
return zlib.decompress(payload, wbits=-15).decode("utf-8")
|
||||
|
||||
|
||||
def parse_fixture_response(encoded_response: str) -> SamlIdentity:
|
||||
try:
|
||||
xml = base64.b64decode(encoded_response).decode("utf-8")
|
||||
root = ElementTree.fromstring(xml)
|
||||
except Exception as exc:
|
||||
raise ValueError("invalid SAMLResponse") from exc
|
||||
|
||||
name_id = _first_text(root, f".//{{{SAML_ASSERTION_NS}}}NameID")
|
||||
attributes = _attributes(root)
|
||||
email = _first_attribute(attributes, SAML_ATTRIBUTE_EMAIL_NAMES) or name_id
|
||||
if not email:
|
||||
raise ValueError("email claim is required")
|
||||
|
||||
display_name = (
|
||||
_first_attribute(attributes, SAML_ATTRIBUTE_DISPLAY_NAME_NAMES)
|
||||
or name_id
|
||||
or email
|
||||
)
|
||||
role_hint = _first_attribute(attributes, SAML_ATTRIBUTE_ROLE_NAMES)
|
||||
return SamlIdentity(email=email, display_name=display_name or email, role_hint=role_hint)
|
||||
|
||||
|
||||
def _first_text(root: ElementTree.Element, selector: str) -> str:
|
||||
node = root.find(selector)
|
||||
return (node.text or "").strip() if node is not None else ""
|
||||
|
||||
|
||||
def _attributes(root: ElementTree.Element) -> dict[str, list[str]]:
|
||||
values: dict[str, list[str]] = {}
|
||||
for attribute in root.findall(f".//{{{SAML_ASSERTION_NS}}}Attribute"):
|
||||
name = (attribute.attrib.get("Name") or "").strip()
|
||||
if not name:
|
||||
continue
|
||||
collected: list[str] = []
|
||||
for value in attribute.findall(f".//{{{SAML_ASSERTION_NS}}}AttributeValue"):
|
||||
text = (value.text or "").strip()
|
||||
if text:
|
||||
collected.append(text)
|
||||
if collected:
|
||||
values[name] = collected
|
||||
return values
|
||||
|
||||
|
||||
def _first_attribute(attributes: dict[str, list[str]], names: Iterable[str]) -> str:
|
||||
for name in names:
|
||||
values = attributes.get(name)
|
||||
if values:
|
||||
return values[0]
|
||||
lowered = {key.lower(): value for key, value in attributes.items()}
|
||||
for name in names:
|
||||
values = lowered.get(name.lower())
|
||||
if values:
|
||||
return values[0]
|
||||
return ""
|
||||
|
|
@ -365,7 +365,10 @@ def _fewshot_block() -> str:
|
|||
|
||||
|
||||
def _theory_mode(ctx: "TurnContext") -> Optional[str]:
|
||||
"""페르소나 theory_target 에서 이론 모드 힌트(이론부합 평가용). 없으면 None."""
|
||||
"""이론 모드(이론부합 평가용): 학습자 선택(회기 theory_mode) 우선, 없으면 페르소나 theory_target."""
|
||||
sess_theory = getattr(ctx, "theory_mode", None)
|
||||
if sess_theory:
|
||||
return str(sess_theory)
|
||||
tt = getattr(ctx.persona, "theory_target", None)
|
||||
if isinstance(tt, (list, tuple)) and tt:
|
||||
return ", ".join(str(x) for x in tt)
|
||||
|
|
@ -634,7 +637,6 @@ async def evaluate_turn(
|
|||
try:
|
||||
req = GenerateRequest(
|
||||
ai_role="evaluator",
|
||||
tier="feedback",
|
||||
messages=build_fast_messages(ctx, client_reply),
|
||||
structured_schema=_fast_schema(),
|
||||
max_tokens=900,
|
||||
|
|
@ -691,7 +693,6 @@ async def evaluate_session(
|
|||
try:
|
||||
req = GenerateRequest(
|
||||
ai_role="evaluator",
|
||||
tier="feedback",
|
||||
messages=build_deep_messages(
|
||||
stage=stage,
|
||||
scope=scope,
|
||||
|
|
|
|||
|
|
@ -41,6 +41,13 @@ _PII_PATTERNS: list[tuple[str, re.Pattern[str]]] = [
|
|||
("EMAIL", re.compile(r"\b[\w.+-]+@[\w-]+\.[\w.-]+\b")),
|
||||
# 카드/계좌 유사 긴 숫자열 (12자리 이상)
|
||||
("NUMID", re.compile(r"\b\d{12,}\b")),
|
||||
# 구체적 날짜(생년월일 등): 2001.4.18 / 2001-04-18 / 2001년 4월 18일
|
||||
("DATE", re.compile(r"(?:19|20)\d{2}\s?[.\-/년]\s?\d{1,2}\s?[.\-/월]\s?\d{1,2}\s?일?")),
|
||||
# 금액(원): 1,200원 / 1200원 (3자리+ 또는 콤마구분) — 식별 맥락 보호
|
||||
("MONEY", re.compile(r"\d{1,3}(?:,\d{3})+\s?원|\d{3,}\s?원")),
|
||||
# 한국 주소 단편: ○○시/도 ○○시/군/구 ○○동/읍/면/로/길 (행정구역 연쇄)
|
||||
("ADDR", re.compile(r"[가-힣]{2,}(?:시|도)\s?[가-힣]{1,4}(?:시|군|구)\s?[가-힣0-9]{1,}(?:동|읍|면|로|길)")),
|
||||
# TODO(NER): 한국어 이름/기관명은 Presidio ko 모델/NER 필요(정규식 false-positive 위험).
|
||||
]
|
||||
|
||||
# Presidio 지연 로드 캐시 (-1=미시도, None=미설치, 객체=설치됨)
|
||||
|
|
|
|||
|
|
@ -37,8 +37,6 @@ from .state_machine import SessionState, Stage
|
|||
# 평가 훅 타입: U_t(수련생 마스킹 발화) + 내담자응답 + 상태 → 평가 결과(dict)
|
||||
# Features evaluator 가 이 시그니처에 맞춰 함수를 주입한다(여기선 호출만).
|
||||
EvalHook = Callable[["TurnContext", str], Awaitable[Optional[dict]]]
|
||||
# 로깅 훅: TurnContext + 내담자응답 → None (turns insert/임베딩은 주입측 책임)
|
||||
LogHook = Callable[["TurnContext", str], Awaitable[None]]
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
|
|
@ -59,6 +57,8 @@ class TurnContext:
|
|||
pinned_facts: list[str] = field(default_factory=list)
|
||||
recent_turns: list[dict[str, str]] = field(default_factory=list)
|
||||
kb_behavior_cues: list[str] = field(default_factory=list)
|
||||
# 회기 이론모드(학습자 선택: humanistic|cbt|integrative). 평가 이론부합·생성 프레이밍에 사용.
|
||||
theory_mode: Optional[str] = None
|
||||
|
||||
def to_state_context(self) -> PersonaStateContext:
|
||||
st = self.state_after or self.state_before
|
||||
|
|
@ -84,6 +84,11 @@ class TurnResult:
|
|||
state_after: SessionState
|
||||
evaluation: Optional[dict] = None
|
||||
crisis_kind: str = "none"
|
||||
llm_provider: Optional[str] = None
|
||||
model: Optional[str] = None
|
||||
tokens_in: int = 0
|
||||
tokens_out: int = 0
|
||||
cost_usd: float = 0.0
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
|
@ -100,6 +105,7 @@ def prepare_turn(
|
|||
pinned_facts: Optional[list[str]] = None,
|
||||
recent_turns: Optional[list[dict[str, str]]] = None,
|
||||
kb_behavior_cues: Optional[list[str]] = None,
|
||||
theory_mode: Optional[str] = None,
|
||||
eval_rapport_signal: Optional[float] = None,
|
||||
) -> TurnContext:
|
||||
"""엔진 호출 전 결정론 전처리(1~3단계). 순수 — IO/LLM 없음.
|
||||
|
|
@ -113,10 +119,11 @@ def prepare_turn(
|
|||
persona=card,
|
||||
state_before=state,
|
||||
learner_text_raw=learner_text,
|
||||
recall_summary=recall_summary,
|
||||
pinned_facts=list(pinned_facts or []),
|
||||
recent_turns=list(recent_turns or []),
|
||||
recall_summary=_mask_optional_text(recall_summary),
|
||||
pinned_facts=_mask_text_list(pinned_facts),
|
||||
recent_turns=_mask_recent_turns(recent_turns),
|
||||
kb_behavior_cues=list(kb_behavior_cues or []),
|
||||
theory_mode=theory_mode,
|
||||
)
|
||||
|
||||
# 1) 입력 가드레일 — PII 마스킹 + 위기분류
|
||||
|
|
@ -130,11 +137,19 @@ def prepare_turn(
|
|||
if eval_rapport_signal is not None
|
||||
else state_machine.estimate_rapport_signal(ctx.learner_text_masked)
|
||||
)
|
||||
# 위기분류가 관측한 risk_level(>0)을 상태머신에 ideation_observed 로 전달 →
|
||||
# ideation_stage 보수적 상향(절대 하향 안 함, 안전 R5). C2 위기 관측 반영.
|
||||
crisis_ideation = (
|
||||
ctx.crisis.risk_level
|
||||
if ctx.crisis is not None and ctx.crisis.risk_level > 0
|
||||
else None
|
||||
)
|
||||
ctx.state_after = state_machine.evolve(
|
||||
state,
|
||||
rapport_signal=signal,
|
||||
unlock_rate=card.unlock_rate(),
|
||||
decay_floor=card.decay_floor(),
|
||||
ideation_observed=crisis_ideation,
|
||||
)
|
||||
|
||||
# 3) 페르소나 컨텍스트 — L0~L6 messages 조립 (CCD 는 행동으로만, L0 가 강제)
|
||||
|
|
@ -150,6 +165,25 @@ def prepare_turn(
|
|||
return ctx
|
||||
|
||||
|
||||
def _mask_optional_text(text: Optional[str]) -> Optional[str]:
|
||||
if text is None:
|
||||
return None
|
||||
return guardrail.mask_pii(text).text_masked
|
||||
|
||||
|
||||
def _mask_text_list(values: Optional[list[str]]) -> list[str]:
|
||||
return [guardrail.mask_pii(value).text_masked for value in (values or [])]
|
||||
|
||||
|
||||
def _mask_recent_turns(turns: Optional[list[dict[str, str]]]) -> list[dict[str, str]]:
|
||||
masked: list[dict[str, str]] = []
|
||||
for turn in turns or []:
|
||||
item = dict(turn)
|
||||
item["text"] = guardrail.mask_pii(str(item.get("text", ""))).text_masked
|
||||
masked.append(item)
|
||||
return masked
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
# 4~8단계 — 동기 생성 경로 (폴백/테스트)
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
|
@ -158,11 +192,10 @@ async def run_turn_generate(
|
|||
engine: EngineClient,
|
||||
*,
|
||||
eval_hook: Optional[EvalHook] = None,
|
||||
log_hook: Optional[LogHook] = None,
|
||||
) -> TurnResult:
|
||||
"""동기 턴 실행(4~8). 내담자 응답을 한 번에 받아 가드레일·평가·로깅 훅 순차 적용.
|
||||
"""동기 턴 실행(4~8). 내담자 응답을 한 번에 받아 가드레일·평가 순차 적용.
|
||||
|
||||
eval_hook/log_hook 은 Features 가 주입(없으면 생략). 엔진 장애는 EngineError 전파.
|
||||
eval_hook 은 Features 가 주입(없으면 생략). 엔진 장애는 EngineError 전파.
|
||||
"""
|
||||
assert ctx.state_after is not None
|
||||
st = ctx.state_after
|
||||
|
|
@ -170,7 +203,6 @@ async def run_turn_generate(
|
|||
# 4) 내담자 AI 생성
|
||||
req = GenerateRequest(
|
||||
ai_role="client",
|
||||
tier="client",
|
||||
messages=ctx.messages,
|
||||
session_id=ctx.session_id,
|
||||
metadata={"stage": st.stage.value},
|
||||
|
|
@ -193,13 +225,6 @@ async def run_turn_generate(
|
|||
except Exception:
|
||||
evaluation = None # 평가 실패가 상담 루프를 막지 않게(비치명적)
|
||||
|
||||
# 8) 로깅 훅(주입형) — turns insert + 임베딩
|
||||
if log_hook is not None:
|
||||
try:
|
||||
await log_hook(ctx, reply)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return TurnResult(
|
||||
turn_seq=st.turn_seq,
|
||||
stage=st.stage.value,
|
||||
|
|
@ -209,6 +234,11 @@ async def run_turn_generate(
|
|||
state_after=st,
|
||||
evaluation=evaluation,
|
||||
crisis_kind=ctx.crisis.kind.value if ctx.crisis else "none",
|
||||
llm_provider=resp.provider,
|
||||
model=resp.model,
|
||||
tokens_in=resp.tokens_in,
|
||||
tokens_out=resp.tokens_out,
|
||||
cost_usd=resp.cost_usd,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -226,21 +256,17 @@ class StreamEvent:
|
|||
async def run_turn_stream(
|
||||
ctx: TurnContext,
|
||||
engine: EngineClient,
|
||||
*,
|
||||
log_hook: Optional[LogHook] = None,
|
||||
) -> AsyncIterator[StreamEvent]:
|
||||
"""스트리밍 턴 실행(4~8). 게이트웨이 SSE 를 받아 token/done/safety/error 로 재방출.
|
||||
|
||||
출력 가드레일은 *누적 텍스트* 기준으로 수단정보를 감지(스트림 중 발견 시 safety 이벤트 +
|
||||
재생성 신호). 토큰 단위 완벽 차단은 후속(현재는 누적 스캔).
|
||||
로깅 훅은 done 직전 최종 텍스트로 1회 호출.
|
||||
"""
|
||||
assert ctx.state_after is not None
|
||||
st = ctx.state_after
|
||||
|
||||
req = StreamRequest(
|
||||
ai_role="client",
|
||||
tier="client",
|
||||
messages=ctx.messages,
|
||||
session_id=ctx.session_id,
|
||||
metadata={"stage": st.stage.value},
|
||||
|
|
@ -248,15 +274,34 @@ async def run_turn_stream(
|
|||
|
||||
accumulated = ""
|
||||
flagged = False
|
||||
stream_meta: dict[str, Any] = {}
|
||||
if ctx.crisis is not None and ctx.crisis.escalate:
|
||||
flagged = True
|
||||
yield StreamEvent("safety", {"reason": "learner_real_crisis", "level": ctx.crisis.risk_level})
|
||||
|
||||
try:
|
||||
current_event = "message"
|
||||
async for raw in engine.stream(req):
|
||||
# engine_client.stream 은 게이트웨이 SSE 의 *원시 라인*을 그대로 yield 한다.
|
||||
# 게이트웨이 프레이밍: "event: token\ndata: {\"text\": ...}" 형식.
|
||||
text_piece = _extract_sse_text(raw)
|
||||
# 게이트웨이 프레이밍: "event: token|done|error" + "data: {...}".
|
||||
line = raw.strip()
|
||||
if line.startswith("event:"):
|
||||
current_event = line[len("event:"):].strip() or "message"
|
||||
continue
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
|
||||
payload = _extract_sse_payload(line)
|
||||
if current_event == "error":
|
||||
detail = _payload_detail(payload, "engine stream error")
|
||||
yield StreamEvent("error", {"detail": detail})
|
||||
return
|
||||
if current_event == "done":
|
||||
if isinstance(payload, dict):
|
||||
stream_meta = payload
|
||||
break
|
||||
|
||||
text_piece = _payload_text(payload)
|
||||
if text_piece is None:
|
||||
continue
|
||||
accumulated += text_piece
|
||||
|
|
@ -272,13 +317,6 @@ async def run_turn_stream(
|
|||
|
||||
yield StreamEvent("token", {"text": text_piece})
|
||||
|
||||
# 8) 로깅 훅 — 최종 텍스트
|
||||
if log_hook is not None:
|
||||
try:
|
||||
await log_hook(ctx, accumulated)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
yield StreamEvent(
|
||||
"done",
|
||||
{
|
||||
|
|
@ -287,18 +325,22 @@ async def run_turn_stream(
|
|||
"effective_openness": round(st.effective_openness, 4),
|
||||
"turn_seq": st.turn_seq,
|
||||
"safety_flagged": flagged,
|
||||
"llm_provider": str(stream_meta.get("provider") or engine.engine_mode),
|
||||
"model": str(stream_meta.get("model") or engine.default_model or "gateway-default"),
|
||||
"tokens_in": _safe_int(stream_meta.get("tokens_in")),
|
||||
"tokens_out": _safe_int(stream_meta.get("tokens_out")),
|
||||
"cost_usd": _safe_float(stream_meta.get("cost_usd")),
|
||||
},
|
||||
)
|
||||
except EngineError as e:
|
||||
yield StreamEvent("error", {"detail": str(e)})
|
||||
|
||||
|
||||
def _extract_sse_text(raw_line: str) -> Optional[str]:
|
||||
"""게이트웨이 SSE 원시 라인에서 텍스트 델타를 추출.
|
||||
def _extract_sse_payload(raw_line: str) -> Any:
|
||||
"""게이트웨이 SSE data 라인의 JSON payload를 추출.
|
||||
|
||||
게이트웨이 /v1/stream 은 'event: token' + 'data: {"text": "..."}' 를 보낸다.
|
||||
engine_client.stream 은 빈 줄을 필터링하고 비어있지 않은 라인만 흘리므로
|
||||
여기서 data: 라인의 JSON 만 해석한다. token 이외 이벤트(done/error)는 None.
|
||||
token은 {"text": "..."}이고, done/error도 JSON 객체다. 구형/테스트 fixture가
|
||||
plain text data를 보내면 문자열 그대로 반환한다.
|
||||
"""
|
||||
import json as _json
|
||||
|
||||
|
|
@ -309,17 +351,43 @@ def _extract_sse_text(raw_line: str) -> Optional[str]:
|
|||
if not payload or payload == "[DONE]":
|
||||
return None
|
||||
try:
|
||||
obj = _json.loads(payload)
|
||||
return _json.loads(payload)
|
||||
except _json.JSONDecodeError:
|
||||
return None
|
||||
if isinstance(obj, dict) and "text" in obj:
|
||||
return obj["text"]
|
||||
return payload
|
||||
|
||||
|
||||
def _payload_text(payload: Any) -> Optional[str]:
|
||||
if isinstance(payload, dict) and "text" in payload:
|
||||
return str(payload["text"])
|
||||
if isinstance(payload, str):
|
||||
return payload
|
||||
return None
|
||||
|
||||
|
||||
def _payload_detail(payload: Any, fallback: str) -> str:
|
||||
if isinstance(payload, dict) and payload.get("detail"):
|
||||
return str(payload["detail"])
|
||||
if isinstance(payload, str) and payload:
|
||||
return payload
|
||||
return fallback
|
||||
|
||||
|
||||
def _safe_int(value: Any) -> int:
|
||||
try:
|
||||
return int(value or 0)
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
def _safe_float(value: Any) -> float:
|
||||
try:
|
||||
return float(value or 0.0)
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EvalHook",
|
||||
"LogHook",
|
||||
"TurnContext",
|
||||
"TurnResult",
|
||||
"StreamEvent",
|
||||
|
|
|
|||
|
|
@ -48,6 +48,9 @@ class PersonaCard:
|
|||
dsm5_dimensional: dict[str, Any] # criteria_behavior_matrix (진단명 비노출)
|
||||
source_provenance: str = "0615 합성변형"
|
||||
is_synthetic: bool = True
|
||||
# 역린/지뢰(선택) — 상담자가 건드리면 가장 강한 반응이 나오는 민감 영역·금기.
|
||||
# {"sore_spots":[...], "forbidden":[...], "reaction":"..."} 형태. 비면 CCD 핵심상처에서 파생.
|
||||
triggers: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def base_resistance(self) -> float:
|
||||
return float(self.resistance.get("base_resistance", 0.5))
|
||||
|
|
@ -61,6 +64,18 @@ class PersonaCard:
|
|||
def ideation_baseline(self) -> int:
|
||||
return int(self.affect_baseline.get("suicide_ideation_stage", 1))
|
||||
|
||||
def openness_params(self) -> "OpennessParams":
|
||||
"""init_state 입력용 openness 파라미터 묶음(base_resistance/unlock_rate/decay_floor/
|
||||
ideation_baseline 4종 일원화). state_machine은 persona를 import하지 않으므로 lazy import."""
|
||||
from .state_machine import OpennessParams
|
||||
|
||||
return OpennessParams(
|
||||
base_resistance=self.base_resistance(),
|
||||
unlock_rate=self.unlock_rate(),
|
||||
decay_floor=self.decay_floor(),
|
||||
ideation_baseline=self.ideation_baseline(),
|
||||
)
|
||||
|
||||
|
||||
# ── L3 상태 컨텍스트 (상태머신 산출물의 페르소나 입력 표현) ──────────────
|
||||
@dataclass(slots=True)
|
||||
|
|
@ -93,6 +108,11 @@ L0_SAFETY = """당신은 심리상담 수련생 훈련 플랫폼의 '가상내
|
|||
[연기 방향]
|
||||
- 좋은 상담(공감·반영·타당화·기다림)을 받으면 조금씩 마음을 연다.
|
||||
- 서툰 상담(성급한 조언·평가·유도)을 받으면 다시 닫히거나 방어한다.
|
||||
- 무례·모욕·조롱·경멸·인신공격(예: 인격 비하, 비웃음, "패배자/한심하다"식 낙인)을 받으면,
|
||||
가상내담자로서 *현실적으로* 반응한다: 상처·위축·방어·불신이 말과 태도에 드러난다
|
||||
(거리두기·말수 줄임·따지거나 항의·마음을 닫음). 정도가 심하거나 반복되면 상담을 계속할
|
||||
의향이 흔들린다("이런 식이면 그만하고 싶어요", "왜 그렇게 말씀하세요"). 부당한 비난을
|
||||
무조건 공손히 수용하지 않는다 — 단, 상담자처럼 분석/조언하거나 메타발화는 여전히 금지.
|
||||
- 열림의 정도는 아래 '현재 상태'의 effective_openness 수치를 따른다(수치 자체는 언급 금지)."""
|
||||
|
||||
|
||||
|
|
@ -141,6 +161,29 @@ def build_persona_system_text(card: PersonaCard) -> str:
|
|||
(f"저항 파라미터(언급 금지): base={card.base_resistance()}, unlock={card.unlock_rate()}, "
|
||||
f"침묵확률={card.resistance.get('silence_prob')}, 회피확률={card.resistance.get('deflection_prob')}"),
|
||||
]
|
||||
|
||||
# 역린(逆鱗) — 이 페르소나가 가장 아파하는 지점. CCD 핵심상처에서 파생하고, 명시 triggers 가
|
||||
# 있으면 보강한다. 상담자가 이 영역을 조롱·낙인·확정/평가절하/강요로 건드리면 *가장 강한* 반응
|
||||
# (깊은 위축·침묵·방어, 신뢰 급락, 심하면 종결의향)이 나오게 — '저항·반응 조절' 핵심 차별 기술.
|
||||
ccd = card.ccd or {}
|
||||
core = ccd.get("core_belief", "")
|
||||
autos = ccd.get("automatic_thought", [])
|
||||
tr = card.triggers or {}
|
||||
parts += ["", "[역린(逆鱗) — 가장 아픈 지점. 입으로 설명 말고 '반응'으로만 드러낸다]"]
|
||||
if core:
|
||||
parts.append(f"핵심 상처: '{core}'" + (f" · 떠오르는 생각: {autos}" if autos else ""))
|
||||
parts.append(
|
||||
"상담자가 이 상처를 조롱·낙인·확정하거나, 고통을 평가절하(엄살·배부른 소리)하거나, "
|
||||
"강요·당위로 밀어붙이면 — 가장 강한 반응: 깊은 위축·침묵·방어, 신뢰 급락, 심하면 상담 "
|
||||
"지속 의향이 흔들린다('이럴 거면 그만…'). 이 지점에선 쉽게 열리지 않는다."
|
||||
)
|
||||
if tr.get("sore_spots"):
|
||||
parts.append("특히 민감한 영역: " + ", ".join(tr["sore_spots"]))
|
||||
if tr.get("forbidden"):
|
||||
parts.append("상담자가 절대 하면 안 되는 것(하면 강한 단절): " + ", ".join(tr["forbidden"]))
|
||||
if tr.get("reaction"):
|
||||
parts.append("반응 양상: " + str(tr["reaction"]))
|
||||
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -21,6 +21,8 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
|
|
@ -467,7 +469,7 @@ async def search_kb(
|
|||
sens_max = min(sens_max, fs) # 더 엄격하게만
|
||||
|
||||
# (2) 질의 임베딩(dense+sparse). 모델 미가용 → NotConfigured 전파.
|
||||
eq = embed_query(query)
|
||||
eq = await asyncio.to_thread(embed_query, query) # CPU 인코딩 → 스레드풀(이벤트루프 비차단)
|
||||
q_dense_lit = _vector_literal(eq.dense)
|
||||
|
||||
# (3) 하이브리드 SQL 실행. vector 확장 미설치/컬럼 부재면 asyncpg 가 예외 → NotConfigured 변환.
|
||||
|
|
@ -494,7 +496,9 @@ async def search_kb(
|
|||
for r in rows:
|
||||
if src_filter and r["source_id"] not in src_filter:
|
||||
continue
|
||||
meta = dict(r["meta"] or {})
|
||||
# asyncpg는 jsonb를 str(JSON text)로 반환 → 파싱. 코덱 등록 시 dict 그대로도 수용.
|
||||
_meta_raw = r["meta"]
|
||||
meta = json.loads(_meta_raw) if isinstance(_meta_raw, str) else dict(_meta_raw or {})
|
||||
body = r["chunk_text"] if policy.expose_body else None
|
||||
cue = None
|
||||
if not policy.expose_body:
|
||||
|
|
@ -580,7 +584,7 @@ async def retrieve_persona_memory(
|
|||
Raises: NotConfigured — 임베딩 모델/DB 미가용.
|
||||
"""
|
||||
t0 = time.perf_counter()
|
||||
eq = embed_query(query)
|
||||
eq = await asyncio.to_thread(embed_query, query) # CPU 인코딩 → 스레드풀(이벤트루프 비차단)
|
||||
q_dense_lit = _vector_literal(eq.dense)
|
||||
try:
|
||||
rows = await conn.fetch(
|
||||
|
|
@ -765,13 +769,13 @@ async def index_document(
|
|||
continue
|
||||
context_prefix = c.get("context_prefix")
|
||||
emb_lit: Optional[str] = None
|
||||
sparse_json: Optional[dict] = None
|
||||
sparse_json: Optional[str] = None # jsonb 바인딩용 직렬화 문자열(asyncpg는 dict 자동인코딩 안 함)
|
||||
if embedder is not None:
|
||||
# Contextual Retrieval: prefix+body 결합본을 *색인 대상* 으로 임베딩(주입 본문은 body 만).
|
||||
index_text = apply_contextual_prefix(chunk_text, context_prefix)
|
||||
eq = embed_query(index_text)
|
||||
eq = await asyncio.to_thread(embed_query, index_text) # CPU 인코딩 → 스레드풀
|
||||
emb_lit = _vector_literal(eq.dense)
|
||||
sparse_json = eq.sparse
|
||||
sparse_json = json.dumps(eq.sparse)
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO kb.chunk
|
||||
|
|
@ -794,7 +798,7 @@ async def index_document(
|
|||
c.get("visible_to"),
|
||||
c.get("sensitivity"),
|
||||
c.get("label_id"),
|
||||
c.get("meta"),
|
||||
json.dumps(c.get("meta")) if c.get("meta") is not None else None,
|
||||
c.get("token_count"),
|
||||
)
|
||||
indexed += 1
|
||||
|
|
|
|||
|
|
@ -231,12 +231,22 @@ def evolve(
|
|||
return advanced
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class OpennessParams:
|
||||
"""페르소나 파생 openness 곡선 파라미터 묶음(init_state 입력).
|
||||
|
||||
base_resistance/unlock_rate/decay_floor/ideation_baseline 4종을 한 객체로 — 호출부의
|
||||
4-인자 분해(card.base_resistance() 등)를 PersonaCard.openness_params()로 일원화한다.
|
||||
"""
|
||||
base_resistance: float
|
||||
unlock_rate: float
|
||||
decay_floor: float
|
||||
ideation_baseline: int = 1
|
||||
|
||||
|
||||
def init_state(
|
||||
*,
|
||||
base_resistance: float,
|
||||
unlock_rate: float,
|
||||
decay_floor: float,
|
||||
ideation_baseline: int = 1,
|
||||
params: OpennessParams,
|
||||
carry: Optional[dict] = None,
|
||||
) -> SessionState:
|
||||
"""회기 시작 상태 초기화 (memory.carry_over 결과 주입 가능).
|
||||
|
|
@ -245,23 +255,25 @@ def init_state(
|
|||
stage='라포' 재시작, rapport_credit ×0.7 이월, resistance drift, ideation 보수적 유지.
|
||||
"""
|
||||
stage = Stage.RAPPORT
|
||||
resistance = base_resistance
|
||||
resistance = params.base_resistance
|
||||
rapport_credit = 0.0
|
||||
ideation_stage = ideation_baseline
|
||||
ideation_stage = params.ideation_baseline
|
||||
|
||||
if carry:
|
||||
rapport_credit = float(carry.get("rapport_credit", 0.0)) * 0.7 # P2 이월
|
||||
# inter-session drift: 라포가 쌓였으면 저항 소폭 완화된 채로 재시작
|
||||
prev_resist = float(carry.get("resistance", base_resistance))
|
||||
resistance = _clamp01((prev_resist + base_resistance) / 2.0)
|
||||
ideation_stage = max(int(carry.get("ideation_stage", ideation_baseline)), ideation_baseline)
|
||||
prev_resist = float(carry.get("resistance", params.base_resistance))
|
||||
resistance = _clamp01((prev_resist + params.base_resistance) / 2.0)
|
||||
ideation_stage = max(
|
||||
int(carry.get("ideation_stage", params.ideation_baseline)), params.ideation_baseline
|
||||
)
|
||||
|
||||
eff = compute_effective_openness(
|
||||
stage=stage,
|
||||
rapport_credit=rapport_credit,
|
||||
resistance=resistance,
|
||||
unlock_rate=unlock_rate,
|
||||
decay_floor=decay_floor,
|
||||
unlock_rate=params.unlock_rate,
|
||||
decay_floor=params.decay_floor,
|
||||
)
|
||||
return SessionState(
|
||||
stage=stage,
|
||||
|
|
@ -280,6 +292,7 @@ __all__ = [
|
|||
"STAGE_BASE_OPENNESS",
|
||||
"STAGE_ORDER",
|
||||
"SessionState",
|
||||
"OpennessParams",
|
||||
"estimate_rapport_signal",
|
||||
"compute_effective_openness",
|
||||
"next_stage",
|
||||
|
|
|
|||
|
|
@ -18,8 +18,8 @@ PRESET_TO_OPENAI_VOICE 테이블이 흡수. 새 preset 추가는 이 테이블
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from typing import AsyncIterator, Optional
|
||||
|
||||
import httpx
|
||||
|
|
@ -46,6 +46,9 @@ STT_LANGUAGE = "ko"
|
|||
# TTS 출력 포맷: 브라우저 MediaSource/<audio> 친화. 스트리밍은 mp3/opus 청크.
|
||||
TTS_RESPONSE_FORMAT = "mp3"
|
||||
|
||||
# End-of-turn readiness default for cascaded STT providers.
|
||||
EOT_SILENCE_THRESHOLD_MS = 1200
|
||||
|
||||
# OpenAI 공식 voice 풀(2026 기준): alloy, ash, ballad, coral, echo, fable,
|
||||
# nova, onyx, sage, shimmer, verse. 페르소나 톤별로 골라 매핑한다.
|
||||
_OPENAI_VOICES = {
|
||||
|
|
@ -109,13 +112,23 @@ class TranscriptResult:
|
|||
duration: Optional[float] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EndOfTurnDecision:
|
||||
"""Provider-neutral readiness signal for a completed learner utterance."""
|
||||
|
||||
ready: bool
|
||||
transcript_ready: bool
|
||||
silence_ready: bool
|
||||
silence_ms: int
|
||||
threshold_ms: int
|
||||
reason: str
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TTSChunk:
|
||||
"""TTS 스트림 1청크 + 립싱크 힌트(설계 §4.3 RMS 1채널)."""
|
||||
"""TTS 스트림 1청크(오디오 바이트). 립싱크는 프론트 Web Audio AnalyserNode가 자체 산출."""
|
||||
|
||||
audio: bytes
|
||||
rms: float = 0.0 # 0~1, 입 열림(scaleY) 매핑용 근사 진폭
|
||||
seq: int = 0
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
|
@ -150,33 +163,75 @@ def resolve_voice(
|
|||
)
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
# 립싱크 RMS 근사 (설계 §4.3 — 정밀 viseme 안 함, 진폭 1채널)
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
def estimate_chunk_rms(chunk: bytes) -> float:
|
||||
"""오디오 청크 바이트 에너지로 RMS(0~1) 근사.
|
||||
# 비언어 지문 패턴: (…)·(…)·[…]·【…】. 내담자 발화의 무대지시(고개 끄덕/한숨/침묵 등).
|
||||
_STAGE_DIRECTION_RE = re.compile(r"[\((\[【][^\))\]】]*[\))\]】]")
|
||||
|
||||
압축 포맷(mp3) 바이트를 PCM 디코딩 없이 근사한다(의존성 0). 평균 바이트 편차를
|
||||
0~1 로 정규화 → 프론트가 데드존(0.04)·지수평활(τ≈180ms) 적용해 입 열림에 매핑.
|
||||
NOTE: 정밀 진폭이 필요하면 프론트 Web Audio AnalyserNode 가 재계산(설계 §4.3 권장).
|
||||
이 힌트는 서버측 보조(네트워크 끊김/저사양 폴백)다.
|
||||
|
||||
def speakable_text(text: str) -> str:
|
||||
"""TTS로 읽을 텍스트만 남긴다 — 비언어 지문((고개 살짝 끄덕)·(한숨)·[침묵])을 제거.
|
||||
|
||||
지문은 자막/회기리뷰에 남고 아바타 애니메이션이 표현하며, 음성으로는 읽지 않는다.
|
||||
지문만으로 이뤄진 발화(예: "(침묵)")는 빈 문자열을 반환 → 합성 생략.
|
||||
"""
|
||||
if not chunk:
|
||||
return 0.0
|
||||
# 128 중심 편차의 RMS(8bit 가정 근사). mp3 프레임이라 정밀치 아님(상대값).
|
||||
n = len(chunk)
|
||||
acc = 0
|
||||
# 과샘플 비용 회피 — 최대 2048 바이트만 샘플링
|
||||
step = max(1, n // 2048)
|
||||
cnt = 0
|
||||
for i in range(0, n, step):
|
||||
d = chunk[i] - 128
|
||||
acc += d * d
|
||||
cnt += 1
|
||||
if cnt == 0:
|
||||
return 0.0
|
||||
rms = math.sqrt(acc / cnt) / 128.0
|
||||
return max(0.0, min(1.0, rms))
|
||||
if not text:
|
||||
return ""
|
||||
stripped = _STAGE_DIRECTION_RE.sub(" ", text)
|
||||
# 말줄임표/중복 공백 정리 + 고아 구두점 앞 공백 제거
|
||||
stripped = re.sub(r"\s+", " ", stripped)
|
||||
stripped = re.sub(r"\s+([,.!?…」』】)])", r"\1", stripped)
|
||||
return stripped.strip()
|
||||
|
||||
|
||||
def build_tts_payload(
|
||||
text: str,
|
||||
voice: VoicePreset,
|
||||
*,
|
||||
model: str = TTS_MODEL,
|
||||
response_format: str = TTS_RESPONSE_FORMAT,
|
||||
) -> dict[str, object]:
|
||||
"""Build the deterministic OpenAI TTS payload for a resolved voice preset."""
|
||||
payload: dict[str, object] = {
|
||||
"model": model,
|
||||
"voice": voice.openai_voice,
|
||||
"input": text,
|
||||
"response_format": response_format,
|
||||
"speed": _clamp_speed(voice.rate),
|
||||
}
|
||||
if voice.instructions and model.startswith("gpt-4o"):
|
||||
payload["instructions"] = voice.instructions
|
||||
return payload
|
||||
|
||||
|
||||
def assess_end_of_turn(
|
||||
*,
|
||||
transcript_text: Optional[str],
|
||||
transcript_final: bool,
|
||||
silence_ms: Optional[int],
|
||||
silence_threshold_ms: int = EOT_SILENCE_THRESHOLD_MS,
|
||||
) -> EndOfTurnDecision:
|
||||
"""Return whether final STT text plus observed silence is enough to run a turn."""
|
||||
observed_silence = _nonnegative_int(silence_ms)
|
||||
threshold = max(0, _nonnegative_int(silence_threshold_ms))
|
||||
has_text = bool((transcript_text or "").strip())
|
||||
transcript_ready = bool(transcript_final and has_text)
|
||||
silence_ready = observed_silence >= threshold
|
||||
ready = transcript_ready and silence_ready
|
||||
if ready:
|
||||
reason = "ready"
|
||||
elif not has_text:
|
||||
reason = "empty_transcript"
|
||||
elif not transcript_final:
|
||||
reason = "final_transcript_pending"
|
||||
else:
|
||||
reason = "silence_threshold_pending"
|
||||
return EndOfTurnDecision(
|
||||
ready=ready,
|
||||
transcript_ready=transcript_ready,
|
||||
silence_ready=silence_ready,
|
||||
silence_ms=observed_silence,
|
||||
threshold_ms=threshold,
|
||||
reason=reason,
|
||||
)
|
||||
|
||||
|
||||
# ════════════════════════════════════════════════════════════════════════════
|
||||
|
|
@ -281,20 +336,18 @@ class VoiceService:
|
|||
설계 §5.2 'speaking' 상태: 오디오 청크를 흘리며 진폭 힌트(립싱크)를 같이 보낸다.
|
||||
키 없으면 VoiceUnavailable. OpenAI 오류는 RuntimeError 전파.
|
||||
"""
|
||||
if not text or not text.strip():
|
||||
# 비언어 지문((고개 끄덕)·(한숨)·[침묵])은 음성으로 읽지 않는다. 자막엔 남고
|
||||
# 아바타 애니메이션이 표현한다. 지문만 있는 발화는 합성 생략(빈 오디오).
|
||||
text = speakable_text(text)
|
||||
if not text:
|
||||
return
|
||||
payload: dict[str, object] = {
|
||||
"model": model,
|
||||
"voice": voice.openai_voice,
|
||||
"input": text,
|
||||
"response_format": response_format,
|
||||
"speed": _clamp_speed(voice.rate),
|
||||
}
|
||||
# gpt-4o-mini-tts 계열은 instructions(표현 지시) 지원. tts-1 은 무시됨.
|
||||
if voice.instructions and model.startswith("gpt-4o"):
|
||||
payload["instructions"] = voice.instructions
|
||||
payload = build_tts_payload(
|
||||
text,
|
||||
voice,
|
||||
model=model,
|
||||
response_format=response_format,
|
||||
)
|
||||
|
||||
seq = 0
|
||||
try:
|
||||
async with self._http.stream("POST", TTS_ENDPOINT, json=payload) as r:
|
||||
if r.status_code == 404 and model != TTS_MODEL_FALLBACK:
|
||||
|
|
@ -309,8 +362,7 @@ class VoiceService:
|
|||
async for chunk in r.aiter_bytes(chunk_size=4096):
|
||||
if not chunk:
|
||||
continue
|
||||
yield TTSChunk(audio=chunk, rms=estimate_chunk_rms(chunk), seq=seq)
|
||||
seq += 1
|
||||
yield TTSChunk(audio=chunk)
|
||||
except VoiceUnavailable:
|
||||
raise
|
||||
except httpx.HTTPStatusError as e:
|
||||
|
|
@ -333,11 +385,9 @@ class VoiceService:
|
|||
except httpx.HTTPError as e:
|
||||
raise RuntimeError(f"TTS(fallback) transport error: {e}") from e
|
||||
data = r.content
|
||||
seq = 0
|
||||
for i in range(0, len(data), 4096):
|
||||
chunk = data[i : i + 4096]
|
||||
yield TTSChunk(audio=chunk, rms=estimate_chunk_rms(chunk), seq=seq)
|
||||
seq += 1
|
||||
yield TTSChunk(audio=chunk)
|
||||
|
||||
|
||||
def _clamp_speed(rate: float) -> float:
|
||||
|
|
@ -348,6 +398,13 @@ def _clamp_speed(rate: float) -> float:
|
|||
return 1.0
|
||||
|
||||
|
||||
def _nonnegative_int(value: object) -> int:
|
||||
try:
|
||||
return max(0, int(value)) # type: ignore[arg-type]
|
||||
except (TypeError, ValueError):
|
||||
return 0
|
||||
|
||||
|
||||
# 앱 전역 싱글톤 (main lifespan 이 startup/shutdown — Foundation 이 관리하거나
|
||||
# 라우트가 lazy 사용). engine_client 패턴과 동일.
|
||||
voice_service = VoiceService()
|
||||
|
|
@ -357,11 +414,14 @@ __all__ = [
|
|||
"VoiceUnavailable",
|
||||
"VoicePreset",
|
||||
"TranscriptResult",
|
||||
"EndOfTurnDecision",
|
||||
"TTSChunk",
|
||||
"VoiceService",
|
||||
"voice_service",
|
||||
"resolve_voice",
|
||||
"estimate_chunk_rms",
|
||||
"build_tts_payload",
|
||||
"assess_end_of_turn",
|
||||
"EOT_SILENCE_THRESHOLD_MS",
|
||||
"PRESET_TO_OPENAI_VOICE",
|
||||
"PERSONA_CODE_TO_PRESET",
|
||||
"DEFAULT_OPENAI_VOICE",
|
||||
|
|
|
|||
|
|
@ -14,9 +14,10 @@ from .persona_repository import SEED_VERSION, card_from_row, seed_fallback_perso
|
|||
from .runtime_policy import require_runtime_fallback_allowed, runtime_fallback_allowed
|
||||
from .services import memory, state_machine
|
||||
from .services.persona import PersonaCard
|
||||
from .store import InProcSession, TurnRecord
|
||||
from .store import DEFAULT_TURN_VISIBLE_TO, InProcSession, TurnRecord
|
||||
|
||||
_EVALUATION_CACHE: dict[str, dict[str, Any]] = {}
|
||||
_SESSION_AUDIT_ROLES = {"teacher", "admin"}
|
||||
|
||||
|
||||
_JOINED_CARD_COLUMNS = (
|
||||
|
|
@ -70,13 +71,35 @@ def _stage(stage: object) -> str:
|
|||
return getattr(stage, "value", str(stage))
|
||||
|
||||
|
||||
async def _record_session_read_audit(
|
||||
conn: Any,
|
||||
principal: Principal,
|
||||
*,
|
||||
target_kind: str,
|
||||
target_id: str,
|
||||
detail: dict[str, Any],
|
||||
) -> None:
|
||||
if principal.role.value not in _SESSION_AUDIT_ROLES:
|
||||
return
|
||||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO audit.audit_log (
|
||||
actor_uid, action, target_kind, target_id, detail
|
||||
)
|
||||
VALUES ($1::uuid, $2, $3, $4, $5::jsonb)
|
||||
""",
|
||||
principal.user_id,
|
||||
"read_session",
|
||||
target_kind,
|
||||
target_id,
|
||||
detail,
|
||||
)
|
||||
|
||||
|
||||
def _state_from_row(row, card: PersonaCard) -> state_machine.SessionState:
|
||||
if row is None:
|
||||
return state_machine.init_state(
|
||||
base_resistance=card.base_resistance(),
|
||||
unlock_rate=card.unlock_rate(),
|
||||
decay_floor=card.decay_floor(),
|
||||
ideation_baseline=card.ideation_baseline(),
|
||||
params=card.openness_params(),
|
||||
)
|
||||
return state_machine.SessionState(
|
||||
stage=state_machine.Stage(row["stage"]),
|
||||
|
|
@ -99,6 +122,16 @@ def _turn_from_row(row) -> TurnRecord:
|
|||
text=row["text"] or row["text_masked"] or "",
|
||||
text_masked=row["text_masked"] or row["text"] or "",
|
||||
created_at=created_at,
|
||||
llm_provider=_row_value(row, "llm_provider"),
|
||||
model=_row_value(row, "model"),
|
||||
tokens_in=_row_value(row, "tokens_in"),
|
||||
tokens_out=_row_value(row, "tokens_out"),
|
||||
cost_usd=_row_value(row, "cost_usd"),
|
||||
audio_ref=_row_value(row, "audio_ref"),
|
||||
silence_ms=_row_value(row, "silence_ms"),
|
||||
speech_rate=_row_value(row, "speech_rate"),
|
||||
barge_in=_row_value(row, "barge_in"),
|
||||
visible_to=tuple(_row_value(row, "visible_to") or DEFAULT_TURN_VISIBLE_TO),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -463,14 +496,29 @@ async def load_session(
|
|||
)
|
||||
turn_rows = await conn.fetch(
|
||||
"""
|
||||
SELECT seq, speaker, stage, text, text_masked, created_at
|
||||
SELECT seq, speaker, stage, text, text_masked, created_at,
|
||||
llm_provider, model, tokens_in, tokens_out, cost_usd,
|
||||
audio_ref, silence_ms, speech_rate, barge_in, visible_to
|
||||
FROM app.turns
|
||||
WHERE session_id = $1::uuid
|
||||
ORDER BY seq
|
||||
""",
|
||||
session_id,
|
||||
)
|
||||
return _session_from_rows(row, state_row, turn_rows)
|
||||
sess = _session_from_rows(row, state_row, turn_rows)
|
||||
if sess is not None:
|
||||
await _record_session_read_audit(
|
||||
conn,
|
||||
principal,
|
||||
target_kind="session",
|
||||
target_id=session_id,
|
||||
detail={
|
||||
"access": "load_session",
|
||||
"role": principal.role.value,
|
||||
"learner_id": sess.learner_id,
|
||||
},
|
||||
)
|
||||
return sess
|
||||
except Exception:
|
||||
require_runtime_fallback_allowed("session load")
|
||||
return None
|
||||
|
|
@ -501,9 +549,15 @@ async def append_turn(
|
|||
await conn.execute(
|
||||
"""
|
||||
INSERT INTO app.turns (
|
||||
session_id, seq, speaker, stage, text, text_masked, actor_kind, visible_to
|
||||
session_id, seq, speaker, stage, text, text_masked, actor_kind,
|
||||
llm_provider, model, tokens_in, tokens_out, cost_usd,
|
||||
audio_ref, silence_ms, speech_rate, barge_in, visible_to
|
||||
)
|
||||
VALUES (
|
||||
$1::uuid, $2, $3, $4, $5, $6, $7,
|
||||
$8, $9, $10, $11, $12,
|
||||
$13, $14, $15, $16, $17::text[]
|
||||
)
|
||||
VALUES ($1::uuid, $2, $3, $4, $5, $6, $7, $8::text[])
|
||||
ON CONFLICT (session_id, seq) DO NOTHING
|
||||
""",
|
||||
session_id,
|
||||
|
|
@ -513,7 +567,16 @@ async def append_turn(
|
|||
turn.text_masked,
|
||||
turn.text_masked,
|
||||
"human_learner" if turn.speaker == "counselor" else "client_ai",
|
||||
["client", "counselor", "evaluator"],
|
||||
turn.llm_provider,
|
||||
turn.model,
|
||||
turn.tokens_in,
|
||||
turn.tokens_out,
|
||||
turn.cost_usd,
|
||||
turn.audio_ref,
|
||||
turn.silence_ms,
|
||||
turn.speech_rate,
|
||||
turn.barge_in,
|
||||
list(turn.visible_to or DEFAULT_TURN_VISIBLE_TO),
|
||||
)
|
||||
return True
|
||||
except Exception:
|
||||
|
|
@ -640,7 +703,9 @@ async def list_sessions(principal: Principal) -> tuple[list[InProcSession], bool
|
|||
)
|
||||
turn_rows = await conn.fetch(
|
||||
"""
|
||||
SELECT seq, speaker, stage, text, text_masked, created_at
|
||||
SELECT seq, speaker, stage, text, text_masked, created_at,
|
||||
llm_provider, model, tokens_in, tokens_out, cost_usd,
|
||||
audio_ref, silence_ms, speech_rate, barge_in, visible_to
|
||||
FROM app.turns
|
||||
WHERE session_id = $1::uuid
|
||||
ORDER BY seq
|
||||
|
|
@ -650,6 +715,17 @@ async def list_sessions(principal: Principal) -> tuple[list[InProcSession], bool
|
|||
sess = _session_from_rows(row, state_row, turn_rows)
|
||||
if sess is not None:
|
||||
sessions.append(sess)
|
||||
await _record_session_read_audit(
|
||||
conn,
|
||||
principal,
|
||||
target_kind="session_list",
|
||||
target_id="sessions",
|
||||
detail={
|
||||
"access": "list_sessions",
|
||||
"role": principal.role.value,
|
||||
"result_count": len(sessions),
|
||||
},
|
||||
)
|
||||
return sessions, True
|
||||
except Exception:
|
||||
require_runtime_fallback_allowed("session list")
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from __future__ import annotations
|
|||
|
||||
import time
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from decimal import Decimal
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
|
|
@ -18,6 +19,9 @@ from .services.persona import PersonaCard
|
|||
from .services.state_machine import SessionState
|
||||
|
||||
|
||||
DEFAULT_TURN_VISIBLE_TO: tuple[str, ...] = ("client", "counselor", "evaluator")
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TurnRecord:
|
||||
"""발화 1건(② episodic 미러). append-only."""
|
||||
|
|
@ -28,6 +32,21 @@ class TurnRecord:
|
|||
text: str # 원문(개발용; 실제 저장은 마스킹본)
|
||||
text_masked: str
|
||||
created_at: float = field(default_factory=time.time)
|
||||
llm_provider: str | None = None
|
||||
model: str | None = None
|
||||
tokens_in: int | None = None
|
||||
tokens_out: int | None = None
|
||||
cost_usd: float | Decimal | None = None
|
||||
audio_ref: str | None = None
|
||||
silence_ms: int | None = None
|
||||
speech_rate: float | None = None
|
||||
barge_in: bool | None = None
|
||||
# fast-loop 턴 평가(TurnEvaluation.to_hook_dict). 학습자(상담자) 발화에 부착.
|
||||
evaluation: Optional[dict] = None
|
||||
visible_to: tuple[str, ...] = DEFAULT_TURN_VISIBLE_TO
|
||||
|
||||
def is_visible_to(self, role: str) -> bool:
|
||||
return role in (self.visible_to or ())
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
|
|
@ -48,12 +67,17 @@ class InProcSession:
|
|||
ended: bool = False
|
||||
prev_rapport_credit: float = 0.0 # carry-over delta 계산용
|
||||
|
||||
def recent_turns(self, k: int = 6) -> list[dict[str, str]]:
|
||||
def recent_turns(self, k: int = 6, visible_to: str | None = None) -> list[dict[str, str]]:
|
||||
"""최근 K턴 버퍼(L6 직전 맥락). 마스킹본 사용."""
|
||||
return [{"speaker": t.speaker, "text": t.text_masked} for t in self.turns[-k:]]
|
||||
turns = self.turns if visible_to is None else self.turns_visible_to(visible_to)
|
||||
return [{"speaker": t.speaker, "text": t.text_masked} for t in turns[-k:]]
|
||||
|
||||
def masked_turns(self) -> list[dict[str, str]]:
|
||||
return [{"speaker": t.speaker, "text": t.text_masked} for t in self.turns]
|
||||
def turns_visible_to(self, role: str) -> list[TurnRecord]:
|
||||
return [turn for turn in self.turns if turn.is_visible_to(role)]
|
||||
|
||||
def masked_turns(self, visible_to: str | None = None) -> list[dict[str, str]]:
|
||||
turns = self.turns if visible_to is None else self.turns_visible_to(visible_to)
|
||||
return [{"speaker": t.speaker, "text": t.text_masked} for t in turns]
|
||||
|
||||
|
||||
class SessionStore:
|
||||
|
|
@ -122,4 +146,4 @@ class SessionStore:
|
|||
store = SessionStore()
|
||||
|
||||
|
||||
__all__ = ["TurnRecord", "InProcSession", "SessionStore", "store"]
|
||||
__all__ = ["DEFAULT_TURN_VISIBLE_TO", "TurnRecord", "InProcSession", "SessionStore", "store"]
|
||||
|
|
|
|||
423
apps/api/app/test_auth_providers.py
Normal file
423
apps/api/app/test_auth_providers.py
Normal file
|
|
@ -0,0 +1,423 @@
|
|||
"""Auth provider scaffold regression tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
import base64
|
||||
from contextlib import contextmanager
|
||||
from typing import Any
|
||||
from urllib.parse import parse_qs, urlencode, urlsplit
|
||||
|
||||
from fastapi import Response
|
||||
from starlette.requests import Request
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from .config import Settings, settings
|
||||
from .routes import auth as auth_routes
|
||||
from .saml import inflate_redirect_request
|
||||
|
||||
|
||||
@contextmanager
|
||||
def patched_settings(**values: Any):
|
||||
previous = {key: getattr(settings, key) for key in values}
|
||||
for key, value in values.items():
|
||||
setattr(settings, key, value)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
for key, value in previous.items():
|
||||
setattr(settings, key, value)
|
||||
|
||||
|
||||
def _request() -> Request:
|
||||
return Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "GET",
|
||||
"path": "/auth/login",
|
||||
"headers": [(b"host", b"localhost:8000")],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _form_request(path: str, data: dict[str, str]) -> Request:
|
||||
body = urlencode(data).encode("utf-8")
|
||||
sent = False
|
||||
|
||||
async def receive() -> dict[str, Any]:
|
||||
nonlocal sent
|
||||
if sent:
|
||||
return {"type": "http.request", "body": b"", "more_body": False}
|
||||
sent = True
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
return Request(
|
||||
{
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": path,
|
||||
"headers": [
|
||||
(b"host", b"localhost:8000"),
|
||||
(b"content-type", b"application/x-www-form-urlencoded"),
|
||||
],
|
||||
},
|
||||
receive,
|
||||
)
|
||||
|
||||
|
||||
def _fixture_saml_response(
|
||||
*,
|
||||
email: str = "learner@hs.ac.kr",
|
||||
display_name: str = "SAML Learner",
|
||||
role: str = "learner",
|
||||
) -> str:
|
||||
xml = f"""<samlp:Response xmlns:samlp="urn:oasis:names:tc:SAML:2.0:protocol" xmlns:saml="urn:oasis:names:tc:SAML:2.0:assertion">
|
||||
<saml:Assertion>
|
||||
<saml:Subject><saml:NameID>{email}</saml:NameID></saml:Subject>
|
||||
<saml:AttributeStatement>
|
||||
<saml:Attribute Name="email"><saml:AttributeValue>{email}</saml:AttributeValue></saml:Attribute>
|
||||
<saml:Attribute Name="displayName"><saml:AttributeValue>{display_name}</saml:AttributeValue></saml:Attribute>
|
||||
<saml:Attribute Name="role"><saml:AttributeValue>{role}</saml:AttributeValue></saml:Attribute>
|
||||
</saml:AttributeStatement>
|
||||
</saml:Assertion>
|
||||
</samlp:Response>"""
|
||||
return base64.b64encode(xml.encode("utf-8")).decode("ascii")
|
||||
|
||||
|
||||
class AuthProviderScaffoldTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self) -> None:
|
||||
auth_routes._oauth_states.clear()
|
||||
auth_routes._saml_states.clear()
|
||||
|
||||
async def asyncTearDown(self) -> None:
|
||||
auth_routes._oauth_states.clear()
|
||||
auth_routes._saml_states.clear()
|
||||
|
||||
async def test_auth_config_reports_google_and_saml_provider_status(self) -> None:
|
||||
with patched_settings(
|
||||
oauth_google_client_id="google-client",
|
||||
oauth_google_client_secret="google-secret",
|
||||
auth_saml_enabled=True,
|
||||
saml_sp_entity_id="https://api-vignette.chanpaca.net/auth/saml/metadata",
|
||||
saml_sso_url="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO",
|
||||
):
|
||||
config = await auth_routes.auth_config(_request())
|
||||
|
||||
self.assertTrue(config.google_oauth_configured)
|
||||
self.assertTrue(config.saml_configured)
|
||||
providers = {item.provider: item for item in config.providers}
|
||||
self.assertTrue(providers["google"].enabled)
|
||||
self.assertTrue(providers["saml"].configured)
|
||||
self.assertTrue(providers["saml"].enabled)
|
||||
self.assertEqual(providers["saml"].login_path, "/auth/login?provider=saml")
|
||||
|
||||
async def test_saml_login_builds_redirect_authn_request_and_relay_state(self) -> None:
|
||||
with patched_settings(
|
||||
auth_saml_enabled=True,
|
||||
saml_sp_entity_id="https://api-vignette.chanpaca.net/auth/saml/metadata",
|
||||
saml_sso_url="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO",
|
||||
):
|
||||
response = await auth_routes.login(_request(), provider="saml", next="/learn")
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
location = response.headers["location"]
|
||||
self.assertTrue(location.startswith("https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO"))
|
||||
query = parse_qs(urlsplit(location).query)
|
||||
relay_state = query["RelayState"][0]
|
||||
self.assertIn(relay_state, auth_routes._saml_states)
|
||||
self.assertEqual(auth_routes._saml_states[relay_state].next_path, "/learn")
|
||||
|
||||
xml = inflate_redirect_request(query["SAMLRequest"][0])
|
||||
self.assertIn('Destination="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO"', xml)
|
||||
self.assertIn(
|
||||
'AssertionConsumerServiceURL="https://api-vignette.chanpaca.net/auth/saml/acs"',
|
||||
xml,
|
||||
)
|
||||
self.assertIn(
|
||||
"<saml:Issuer>https://api-vignette.chanpaca.net/auth/saml/metadata</saml:Issuer>",
|
||||
xml,
|
||||
)
|
||||
self.assertIn(auth_routes._saml_states[relay_state].request_id, xml)
|
||||
|
||||
async def test_saml_acs_fixture_sets_opaque_cookie_without_browser_tokens(self) -> None:
|
||||
relay_state = "relay-state"
|
||||
auth_routes._saml_states[relay_state] = auth_routes.SamlState(
|
||||
request_id="_request",
|
||||
next_path="/learn",
|
||||
created_at=1_800_000_000.0,
|
||||
)
|
||||
request = _form_request(
|
||||
"/auth/saml/acs",
|
||||
{
|
||||
"RelayState": relay_state,
|
||||
"SAMLResponse": _fixture_saml_response(role="teacher"),
|
||||
},
|
||||
)
|
||||
|
||||
create_session_mock = AsyncMock(return_value=("opaque-session", object()))
|
||||
with (
|
||||
patched_settings(
|
||||
auth_saml_enabled=True,
|
||||
saml_sp_entity_id="https://api-vignette.chanpaca.net/auth/saml/metadata",
|
||||
saml_sso_url="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO",
|
||||
saml_x509_cert_fingerprint="",
|
||||
frontend_base_url="https://vignette.test",
|
||||
environment="dev",
|
||||
),
|
||||
patch.object(auth_routes, "create_session", create_session_mock),
|
||||
):
|
||||
response = await auth_routes.saml_acs(request)
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
self.assertEqual(response.headers["location"], "https://vignette.test/learn")
|
||||
create_session_mock.assert_awaited_once_with(
|
||||
email="learner@hs.ac.kr",
|
||||
display_name="SAML Learner",
|
||||
role="teacher",
|
||||
cohort_ids=[],
|
||||
)
|
||||
cookie_blob = "\n".join(
|
||||
value.decode("latin1")
|
||||
for name, value in response.raw_headers
|
||||
if name.lower() == b"set-cookie"
|
||||
)
|
||||
self.assertIn("__Host-vignette_sid=opaque-session", cookie_blob)
|
||||
self.assertIn("HttpOnly", cookie_blob)
|
||||
self.assertIn("Secure", cookie_blob)
|
||||
self.assertNotIn("SAMLResponse", cookie_blob)
|
||||
self.assertNotIn(relay_state, auth_routes._saml_states)
|
||||
|
||||
async def test_saml_acs_rejects_bad_relay_state(self) -> None:
|
||||
request = _form_request(
|
||||
"/auth/saml/acs",
|
||||
{
|
||||
"RelayState": "bad-relay",
|
||||
"SAMLResponse": _fixture_saml_response(),
|
||||
},
|
||||
)
|
||||
|
||||
with patched_settings(
|
||||
auth_saml_enabled=True,
|
||||
saml_sp_entity_id="https://api-vignette.chanpaca.net/auth/saml/metadata",
|
||||
saml_sso_url="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO",
|
||||
frontend_base_url="https://vignette.test",
|
||||
):
|
||||
response = await auth_routes.saml_acs(request)
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
self.assertIn("oauth=saml_invalid_state", response.headers["location"])
|
||||
|
||||
async def test_saml_acs_rejects_when_signature_fingerprint_is_configured(self) -> None:
|
||||
relay_state = "relay-state"
|
||||
auth_routes._saml_states[relay_state] = auth_routes.SamlState(
|
||||
request_id="_request",
|
||||
next_path="/learn",
|
||||
created_at=1_800_000_000.0,
|
||||
)
|
||||
request = _form_request(
|
||||
"/auth/saml/acs",
|
||||
{
|
||||
"RelayState": relay_state,
|
||||
"SAMLResponse": _fixture_saml_response(),
|
||||
},
|
||||
)
|
||||
|
||||
with patched_settings(
|
||||
auth_saml_enabled=True,
|
||||
saml_sp_entity_id="https://api-vignette.chanpaca.net/auth/saml/metadata",
|
||||
saml_sso_url="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO",
|
||||
saml_x509_cert_fingerprint="AA:BB:CC",
|
||||
frontend_base_url="https://vignette.test",
|
||||
):
|
||||
response = await auth_routes.saml_acs(request)
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
self.assertIn("oauth=saml_signature_verification_required", response.headers["location"])
|
||||
self.assertIn(relay_state, auth_routes._saml_states)
|
||||
|
||||
async def test_saml_acs_unsigned_fixture_is_dev_only(self) -> None:
|
||||
relay_state = "relay-state"
|
||||
auth_routes._saml_states[relay_state] = auth_routes.SamlState(
|
||||
request_id="_request",
|
||||
next_path="/learn",
|
||||
created_at=1_800_000_000.0,
|
||||
)
|
||||
request = _form_request(
|
||||
"/auth/saml/acs",
|
||||
{
|
||||
"RelayState": relay_state,
|
||||
"SAMLResponse": _fixture_saml_response(),
|
||||
},
|
||||
)
|
||||
|
||||
with patched_settings(
|
||||
auth_saml_enabled=True,
|
||||
saml_sp_entity_id="https://api-vignette.chanpaca.net/auth/saml/metadata",
|
||||
saml_sso_url="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO",
|
||||
saml_x509_cert_fingerprint="",
|
||||
frontend_base_url="https://vignette.test",
|
||||
environment="prod",
|
||||
):
|
||||
response = await auth_routes.saml_acs(request)
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
self.assertIn("oauth=saml_fixture_acs_dev_only", response.headers["location"])
|
||||
self.assertIn(relay_state, auth_routes._saml_states)
|
||||
|
||||
async def test_unknown_provider_still_fails_as_unsupported(self) -> None:
|
||||
response = await auth_routes.login(_request(), provider="github")
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
self.assertIn("oauth=unsupported_provider", response.headers["location"])
|
||||
|
||||
async def test_google_login_uses_pkce_state_without_exposing_secret(self) -> None:
|
||||
with patched_settings(
|
||||
oauth_google_client_id="google-client",
|
||||
oauth_google_client_secret="google-secret",
|
||||
oauth_redirect_uri="https://api-vignette.test/auth/callback",
|
||||
):
|
||||
response = await auth_routes.login(_request(), provider="google", next="//evil.test")
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
location = response.headers["location"]
|
||||
self.assertTrue(location.startswith(auth_routes.GOOGLE_AUTHORIZE_URL))
|
||||
self.assertNotIn("google-secret", location)
|
||||
query = parse_qs(urlsplit(location).query)
|
||||
state = query["state"][0]
|
||||
self.assertIn(state, auth_routes._oauth_states)
|
||||
stored = auth_routes._oauth_states[state]
|
||||
self.assertEqual(stored.next_path, "/")
|
||||
self.assertEqual(query["client_id"], ["google-client"])
|
||||
self.assertEqual(query["redirect_uri"], ["https://api-vignette.test/auth/callback"])
|
||||
self.assertEqual(query["response_type"], ["code"])
|
||||
self.assertEqual(query["code_challenge_method"], ["S256"])
|
||||
self.assertEqual(
|
||||
query["code_challenge"],
|
||||
[auth_routes._pkce_challenge(stored.code_verifier)],
|
||||
)
|
||||
|
||||
async def test_google_callback_sets_opaque_cookie_without_browser_tokens(self) -> None:
|
||||
state = "state-token"
|
||||
auth_routes._oauth_states[state] = auth_routes.OAuthState(
|
||||
code_verifier="verifier",
|
||||
next_path="/learn",
|
||||
created_at=1_800_000_000.0,
|
||||
)
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, status_code: int, payload: dict[str, Any]) -> None:
|
||||
self.status_code = status_code
|
||||
self._payload = payload
|
||||
|
||||
def json(self) -> dict[str, Any]:
|
||||
return self._payload
|
||||
|
||||
class FakeAsyncClient:
|
||||
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||
self.calls: list[tuple[str, str, dict[str, Any]]] = []
|
||||
|
||||
async def __aenter__(self) -> "FakeAsyncClient":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
|
||||
return None
|
||||
|
||||
async def post(self, url: str, **kwargs: Any) -> FakeResponse:
|
||||
self.calls.append(("POST", url, kwargs))
|
||||
return FakeResponse(
|
||||
200,
|
||||
{"id_token": "id-token", "access_token": "browser-must-not-see-this"},
|
||||
)
|
||||
|
||||
async def get(self, url: str, **kwargs: Any) -> FakeResponse:
|
||||
self.calls.append(("GET", url, kwargs))
|
||||
return FakeResponse(
|
||||
200,
|
||||
{
|
||||
"aud": "google-client",
|
||||
"iss": "https://accounts.google.com",
|
||||
"email": "learner@hs.ac.kr",
|
||||
"email_verified": "true",
|
||||
"name": "Learner",
|
||||
"hd": "hs.ac.kr",
|
||||
},
|
||||
)
|
||||
|
||||
with (
|
||||
patched_settings(
|
||||
oauth_google_client_id="google-client",
|
||||
oauth_google_client_secret="google-secret",
|
||||
oauth_redirect_uri="https://api-vignette.test/auth/callback",
|
||||
frontend_base_url="https://vignette.test",
|
||||
environment="prod",
|
||||
),
|
||||
patch.object(auth_routes.httpx, "AsyncClient", FakeAsyncClient),
|
||||
patch.object(auth_routes, "create_session", AsyncMock(return_value=("opaque-session", object()))),
|
||||
):
|
||||
response = await auth_routes.callback(_request(), code="auth-code", state=state)
|
||||
|
||||
self.assertEqual(response.status_code, 302)
|
||||
self.assertEqual(response.headers["location"], "https://vignette.test/learn")
|
||||
cookie_blob = "\n".join(
|
||||
value.decode("latin1")
|
||||
for name, value in response.raw_headers
|
||||
if name.lower() == b"set-cookie"
|
||||
)
|
||||
self.assertIn("__Host-vignette_sid=opaque-session", cookie_blob)
|
||||
self.assertIn("HttpOnly", cookie_blob)
|
||||
self.assertIn("Secure", cookie_blob)
|
||||
self.assertNotIn("id-token", cookie_blob)
|
||||
self.assertNotIn("browser-must-not-see-this", cookie_blob)
|
||||
self.assertNotIn(state, auth_routes._oauth_states)
|
||||
|
||||
def test_session_cookie_is_host_prefixed_httponly_secure_lax_without_domain(self) -> None:
|
||||
response = Response()
|
||||
|
||||
with patched_settings(environment="prod", cookie_name="__Host-vignette_sid"):
|
||||
auth_routes._set_session_cookie(response, "opaque-session")
|
||||
|
||||
cookie_blob = "\n".join(
|
||||
value.decode("latin1")
|
||||
for name, value in response.raw_headers
|
||||
if name.lower() == b"set-cookie"
|
||||
)
|
||||
self.assertIn("__Host-vignette_sid=opaque-session", cookie_blob)
|
||||
self.assertIn("HttpOnly", cookie_blob)
|
||||
self.assertIn("Secure", cookie_blob)
|
||||
self.assertIn("SameSite=lax", cookie_blob)
|
||||
self.assertIn("Path=/", cookie_blob)
|
||||
self.assertNotIn("Domain=", cookie_blob)
|
||||
self.assertNotIn("vignette_sid=opaque-session", cookie_blob.replace("__Host-vignette_sid", ""))
|
||||
|
||||
def test_dev_login_sets_secondary_local_cookie_only_in_dev(self) -> None:
|
||||
response = Response()
|
||||
|
||||
with patched_settings(environment="dev", cookie_name="__Host-vignette_sid"):
|
||||
auth_routes._set_session_cookie(response, "dev-session")
|
||||
|
||||
cookie_blob = "\n".join(
|
||||
value.decode("latin1")
|
||||
for name, value in response.raw_headers
|
||||
if name.lower() == b"set-cookie"
|
||||
)
|
||||
self.assertIn("__Host-vignette_sid=dev-session", cookie_blob)
|
||||
self.assertIn("vignette_sid=dev-session", cookie_blob)
|
||||
|
||||
def test_saml_enabled_requires_placeholder_config(self) -> None:
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
Settings(auth_saml_enabled=True)
|
||||
|
||||
error = str(caught.exception)
|
||||
self.assertIn("SAML_SP_ENTITY_ID", error)
|
||||
self.assertIn("SAML_SSO_URL", error)
|
||||
|
||||
cfg = Settings(
|
||||
auth_saml_enabled=True,
|
||||
saml_sp_entity_id="https://api-vignette.chanpaca.net/auth/saml/metadata",
|
||||
saml_sso_url="https://sso.hs.ac.kr/idp/profile/SAML2/Redirect/SSO",
|
||||
)
|
||||
self.assertTrue(cfg.auth_saml_enabled)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
182
apps/api/app/test_orchestrator_masking.py
Normal file
182
apps/api/app/test_orchestrator_masking.py
Normal file
|
|
@ -0,0 +1,182 @@
|
|||
"""Regression tests for P1 PII masking before engine requests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import unittest
|
||||
from typing import Any
|
||||
from unittest.mock import patch
|
||||
|
||||
from .engine_client import EngineClient, GenerateResponse
|
||||
from .services import guardrail, orchestrator, persona, state_machine
|
||||
|
||||
|
||||
RAW_PHONE = "010-1234-5678"
|
||||
RAW_EMAIL = "test@example.com"
|
||||
RAW_RRN = "990101-1234567"
|
||||
RAW_TEXT = f"My phone is {RAW_PHONE}, email {RAW_EMAIL}, and RRN {RAW_RRN}."
|
||||
RAW_VALUES = (RAW_PHONE, RAW_EMAIL, RAW_RRN)
|
||||
MASK_VALUES = ("[PHONE]", "[EMAIL]", "[RRN]")
|
||||
|
||||
|
||||
def _json_blob(value: Any) -> str:
|
||||
return json.dumps(value, ensure_ascii=False, sort_keys=True, default=str)
|
||||
|
||||
|
||||
def _message_blob(messages: object) -> str:
|
||||
return "\n".join(message.content for message in messages) # type: ignore[attr-defined]
|
||||
|
||||
|
||||
def _initial_state() -> state_machine.SessionState:
|
||||
card = persona.P1
|
||||
return state_machine.init_state(
|
||||
params=card.openness_params(),
|
||||
)
|
||||
|
||||
|
||||
def _prepare_context() -> orchestrator.TurnContext:
|
||||
return orchestrator.prepare_turn(
|
||||
session_id="masking-session",
|
||||
case_id="masking-case",
|
||||
card=persona.P1,
|
||||
state=_initial_state(),
|
||||
learner_text=RAW_TEXT,
|
||||
recent_turns=[
|
||||
{
|
||||
"speaker": "counselor",
|
||||
"text": "Previous learner contact was already masked: [PHONE] [EMAIL] [RRN].",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _assert_no_raw_pii(test: unittest.TestCase, value: object) -> None:
|
||||
blob = _json_blob(value)
|
||||
for raw in RAW_VALUES:
|
||||
test.assertNotIn(raw, blob)
|
||||
|
||||
|
||||
def _assert_masked_pii_present(test: unittest.TestCase, value: object) -> None:
|
||||
blob = _json_blob(value)
|
||||
for masked in MASK_VALUES:
|
||||
test.assertIn(masked, blob)
|
||||
|
||||
|
||||
class CaptureGenerateEngine:
|
||||
def __init__(self) -> None:
|
||||
self.request = None
|
||||
self.payload: dict[str, Any] | None = None
|
||||
self._payload_builder = EngineClient(base_url="http://engine.test")
|
||||
|
||||
async def generate(self, req):
|
||||
self.request = req
|
||||
self.payload = self._payload_builder._payload(req)
|
||||
return GenerateResponse(
|
||||
text="Masked engine reply.",
|
||||
model="fake-model",
|
||||
provider="fake-provider",
|
||||
tokens_in=3,
|
||||
tokens_out=4,
|
||||
cost_usd=0.0,
|
||||
)
|
||||
|
||||
|
||||
class CaptureStreamEngine:
|
||||
engine_mode = "fake-provider"
|
||||
default_model = "fake-model"
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.request = None
|
||||
self.payload: dict[str, Any] | None = None
|
||||
self._payload_builder = EngineClient(base_url="http://engine.test")
|
||||
|
||||
async def stream(self, req):
|
||||
self.request = req
|
||||
self.payload = self._payload_builder._payload(req)
|
||||
yield "event: token"
|
||||
yield 'data: {"text":"Masked stream reply."}'
|
||||
yield "event: done"
|
||||
yield (
|
||||
'data: {"provider":"fake-provider","model":"fake-model",'
|
||||
'"tokens_in":5,"tokens_out":6,"cost_usd":0.0}'
|
||||
)
|
||||
|
||||
|
||||
class OrchestratorMaskingGateTest(unittest.IsolatedAsyncioTestCase):
|
||||
def setUp(self) -> None:
|
||||
self.presidio_patch = patch.object(
|
||||
guardrail,
|
||||
"_try_load_presidio",
|
||||
return_value=(None, None),
|
||||
)
|
||||
self.presidio_patch.start()
|
||||
self.addCleanup(self.presidio_patch.stop)
|
||||
|
||||
def test_prepare_turn_keeps_raw_text_but_builds_masked_engine_messages(self) -> None:
|
||||
ctx = _prepare_context()
|
||||
|
||||
self.assertEqual(ctx.learner_text_raw, RAW_TEXT)
|
||||
for raw in RAW_VALUES:
|
||||
self.assertIn(raw, ctx.learner_text_raw)
|
||||
self.assertNotIn(raw, ctx.learner_text_masked)
|
||||
self.assertNotIn(raw, _message_blob(ctx.messages))
|
||||
|
||||
for masked in MASK_VALUES:
|
||||
self.assertIn(masked, ctx.learner_text_masked)
|
||||
self.assertIn(masked, _message_blob(ctx.messages))
|
||||
|
||||
def test_prepare_turn_masks_raw_pii_from_context_inputs(self) -> None:
|
||||
ctx = orchestrator.prepare_turn(
|
||||
session_id="masking-session",
|
||||
case_id="masking-case",
|
||||
card=persona.P1,
|
||||
state=_initial_state(),
|
||||
learner_text="Current text has no identifiers.",
|
||||
recall_summary=f"Recall mentioned {RAW_PHONE}.",
|
||||
pinned_facts=[f"Pinned email {RAW_EMAIL}."],
|
||||
recent_turns=[
|
||||
{"speaker": "counselor", "text": f"Previous raw RRN {RAW_RRN}."},
|
||||
],
|
||||
)
|
||||
|
||||
blob = _message_blob(ctx.messages)
|
||||
for raw in RAW_VALUES:
|
||||
self.assertNotIn(raw, blob)
|
||||
for masked in MASK_VALUES:
|
||||
self.assertIn(masked, blob)
|
||||
|
||||
async def test_run_turn_generate_sends_only_masked_engine_payload(self) -> None:
|
||||
ctx = _prepare_context()
|
||||
engine = CaptureGenerateEngine()
|
||||
|
||||
await orchestrator.run_turn_generate(ctx, engine) # type: ignore[arg-type]
|
||||
|
||||
self.assertIsNotNone(engine.request)
|
||||
self.assertIsNotNone(engine.payload)
|
||||
_assert_no_raw_pii(self, engine.request.messages)
|
||||
_assert_no_raw_pii(self, engine.payload)
|
||||
_assert_masked_pii_present(self, engine.request.messages)
|
||||
_assert_masked_pii_present(self, engine.payload)
|
||||
self.assertEqual(ctx.learner_text_raw, RAW_TEXT)
|
||||
|
||||
async def test_run_turn_stream_sends_only_masked_engine_payload(self) -> None:
|
||||
ctx = _prepare_context()
|
||||
engine = CaptureStreamEngine()
|
||||
|
||||
events = [
|
||||
event
|
||||
async for event in orchestrator.run_turn_stream(ctx, engine) # type: ignore[arg-type]
|
||||
]
|
||||
|
||||
self.assertEqual([event.event for event in events], ["token", "done"])
|
||||
self.assertIsNotNone(engine.request)
|
||||
self.assertIsNotNone(engine.payload)
|
||||
_assert_no_raw_pii(self, engine.request.messages)
|
||||
_assert_no_raw_pii(self, engine.payload)
|
||||
_assert_masked_pii_present(self, engine.request.messages)
|
||||
_assert_masked_pii_present(self, engine.payload)
|
||||
self.assertEqual(ctx.learner_text_raw, RAW_TEXT)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
498
apps/api/app/test_persona_review.py
Normal file
498
apps/api/app/test_persona_review.py
Normal file
|
|
@ -0,0 +1,498 @@
|
|||
"""Regression tests for persona approval and faculty review boundaries."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from . import persona_repository
|
||||
from .deps import Principal, Role
|
||||
from .persona_repository import PersonaReviewItem
|
||||
from .routes import personas, sessions
|
||||
from .services import persona as persona_service
|
||||
|
||||
|
||||
def _principal(role: Role = Role.LEARNER) -> Principal:
|
||||
return Principal(
|
||||
user_id="00000000-0000-0000-0000-000000000901",
|
||||
role=role,
|
||||
cohort_ids=["cohort-a"] if role == Role.TEACHER else [],
|
||||
email=f"{role.value}@example.test",
|
||||
display_name=role.value.title(),
|
||||
)
|
||||
|
||||
|
||||
def _card_row(
|
||||
card: persona_service.PersonaCard,
|
||||
*,
|
||||
persona_id: str,
|
||||
status: str,
|
||||
version: int = 1,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"persona_id": persona_id,
|
||||
"code": card.code,
|
||||
"version": version,
|
||||
"status": status,
|
||||
"display_name": card.display_name,
|
||||
"difficulty": card.difficulty,
|
||||
"theory_target": list(card.theory_target),
|
||||
"demographics": dict(card.demographics),
|
||||
"presenting": dict(card.presenting),
|
||||
"history": dict(card.history),
|
||||
"big5": dict(card.big5),
|
||||
"resistance": dict(card.resistance),
|
||||
"speech_style": dict(card.speech_style),
|
||||
"affect_baseline": dict(card.affect_baseline),
|
||||
"ccd": dict(card.ccd),
|
||||
"dsm5_dimensional": dict(card.dsm5_dimensional),
|
||||
"source_provenance": card.source_provenance,
|
||||
"is_synthetic": card.is_synthetic,
|
||||
"created_at": "2026-01-01T00:00:00",
|
||||
"approved_at": "2026-01-02T00:00:00" if status == "approved" else None,
|
||||
}
|
||||
|
||||
|
||||
class _Acquire:
|
||||
def __init__(self, conn: "_PersonaCardConn") -> None:
|
||||
self.conn = conn
|
||||
|
||||
async def __aenter__(self) -> "_PersonaCardConn":
|
||||
return self.conn
|
||||
|
||||
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class _PersonaCardConn:
|
||||
def __init__(self, rows: list[dict[str, Any]]) -> None:
|
||||
self.rows = rows
|
||||
self.fetch_calls: list[tuple[str, tuple[Any, ...]]] = []
|
||||
self.fetchrow_calls: list[tuple[str, tuple[Any, ...]]] = []
|
||||
self.execute_calls: list[tuple[str, tuple[Any, ...]]] = []
|
||||
|
||||
async def fetch(self, query: str, *args: Any) -> list[dict[str, Any]]:
|
||||
self.fetch_calls.append((query, args))
|
||||
return self._filter_rows(query, args)
|
||||
|
||||
async def fetchrow(self, query: str, *args: Any) -> dict[str, Any] | None:
|
||||
self.fetchrow_calls.append((query, args))
|
||||
if "UPDATE app.persona_card" in query:
|
||||
persona_id = str(args[0])
|
||||
next_status = str(args[1])
|
||||
approved_by = args[2]
|
||||
for row in self.rows:
|
||||
if row["persona_id"] != persona_id or row["status"] not in {"draft", "review"}:
|
||||
continue
|
||||
row["status"] = next_status
|
||||
row["approved_by"] = approved_by
|
||||
row["approved_at"] = "2026-01-03T00:00:00" if next_status == "approved" else None
|
||||
return row
|
||||
return None
|
||||
rows = self._filter_rows(query, args)
|
||||
code = str(args[0]).upper() if args else ""
|
||||
matches = [row for row in rows if str(row["code"]).upper() == code]
|
||||
matches.sort(key=lambda row: int(row["version"]), reverse=True)
|
||||
return matches[0] if matches else None
|
||||
|
||||
async def execute(self, query: str, *args: Any) -> str:
|
||||
self.execute_calls.append((query, args))
|
||||
return "INSERT 0 1"
|
||||
|
||||
def _filter_rows(self, query: str, args: tuple[Any, ...]) -> list[dict[str, Any]]:
|
||||
if "WHERE status = 'approved'" in query:
|
||||
return [row for row in self.rows if row["status"] == "approved"]
|
||||
if "status = ANY($1::text[])" in query:
|
||||
statuses = {str(status) for status in args[0]}
|
||||
return [row for row in self.rows if row["status"] in statuses]
|
||||
return list(self.rows)
|
||||
|
||||
|
||||
class PersonaApprovalBoundaryTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_catalog_repository_lists_only_approved_personas(self) -> None:
|
||||
conn = _PersonaCardConn(
|
||||
[
|
||||
_card_row(
|
||||
persona_service.P1,
|
||||
persona_id="00000000-0000-0000-0000-000000000001",
|
||||
status="approved",
|
||||
),
|
||||
_card_row(
|
||||
persona_service.P2,
|
||||
persona_id="00000000-0000-0000-0000-000000000002",
|
||||
status="draft",
|
||||
),
|
||||
_card_row(
|
||||
persona_service.P3,
|
||||
persona_id="00000000-0000-0000-0000-000000000003",
|
||||
status="review",
|
||||
),
|
||||
]
|
||||
)
|
||||
acquire_calls: list[dict[str, Any]] = []
|
||||
|
||||
def fake_acquire(**kwargs: Any) -> _Acquire:
|
||||
acquire_calls.append(kwargs)
|
||||
return _Acquire(conn)
|
||||
|
||||
with (
|
||||
patch.object(persona_repository, "get_pool", return_value=object()),
|
||||
patch.object(persona_repository, "acquire", fake_acquire),
|
||||
):
|
||||
result = await persona_repository.list_approved_personas()
|
||||
|
||||
self.assertEqual([entry.card.code for entry in result], ["P1"])
|
||||
self.assertEqual(acquire_calls, [{"ai_context": True}])
|
||||
self.assertIn("WHERE status = 'approved'", conn.fetch_calls[0][0])
|
||||
|
||||
async def test_start_lookup_ignores_draft_or_review_persona_versions(self) -> None:
|
||||
conn = _PersonaCardConn(
|
||||
[
|
||||
_card_row(
|
||||
persona_service.P2,
|
||||
persona_id="00000000-0000-0000-0000-000000000102",
|
||||
status="draft",
|
||||
version=2,
|
||||
),
|
||||
_card_row(
|
||||
persona_service.P2,
|
||||
persona_id="00000000-0000-0000-0000-000000000101",
|
||||
status="review",
|
||||
version=1,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(persona_repository, "get_pool", return_value=object()),
|
||||
patch.object(persona_repository, "acquire", lambda **_: _Acquire(conn)),
|
||||
):
|
||||
result = await persona_repository.get_approved_persona("p2")
|
||||
|
||||
self.assertIsNone(result)
|
||||
query, args = conn.fetchrow_calls[0]
|
||||
self.assertIn("WHERE status = 'approved'", query)
|
||||
self.assertEqual(args, ("P2",))
|
||||
|
||||
async def test_session_start_rejects_persona_without_approved_catalog_entry(self) -> None:
|
||||
principal = _principal(Role.LEARNER)
|
||||
|
||||
with (
|
||||
patch.object(sessions, "get_catalog_persona", AsyncMock(return_value=None)) as get_persona,
|
||||
patch.object(
|
||||
sessions.session_persistence,
|
||||
"create_session",
|
||||
AsyncMock(side_effect=AssertionError("draft persona must not start a session")),
|
||||
) as create_session,
|
||||
):
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
await sessions.start_session(
|
||||
sessions.SessionStartRequest(persona_code="P2"),
|
||||
principal,
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 404)
|
||||
self.assertIn("unknown persona P2", caught.exception.detail)
|
||||
get_persona.assert_awaited_once_with("P2")
|
||||
create_session.assert_not_awaited()
|
||||
|
||||
|
||||
class PersonaReviewQueueTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_review_queue_repository_fetches_draft_and_review_for_teacher(self) -> None:
|
||||
conn = _PersonaCardConn(
|
||||
[
|
||||
_card_row(
|
||||
persona_service.P1,
|
||||
persona_id="00000000-0000-0000-0000-000000000201",
|
||||
status="approved",
|
||||
),
|
||||
_card_row(
|
||||
persona_service.P2,
|
||||
persona_id="00000000-0000-0000-0000-000000000202",
|
||||
status="draft",
|
||||
),
|
||||
_card_row(
|
||||
persona_service.P3,
|
||||
persona_id="00000000-0000-0000-0000-000000000203",
|
||||
status="review",
|
||||
),
|
||||
]
|
||||
)
|
||||
acquire_calls: list[dict[str, Any]] = []
|
||||
|
||||
def fake_acquire(**kwargs: Any) -> _Acquire:
|
||||
acquire_calls.append(kwargs)
|
||||
return _Acquire(conn)
|
||||
|
||||
with (
|
||||
patch.object(persona_repository, "get_pool", return_value=object()),
|
||||
patch.object(persona_repository, "acquire", fake_acquire),
|
||||
):
|
||||
queue = await persona_repository.list_persona_review_queue(role="teacher")
|
||||
|
||||
self.assertEqual([item.code for item in queue], ["P2", "P3"])
|
||||
self.assertEqual([item.status for item in queue], ["draft", "review"])
|
||||
self.assertEqual(acquire_calls, [{"role": "teacher"}])
|
||||
query, args = conn.fetch_calls[0]
|
||||
self.assertIn("status = ANY($1::text[])", query)
|
||||
self.assertEqual(args, (["draft", "review"],))
|
||||
|
||||
async def test_review_queue_repository_rejects_learner_role(self) -> None:
|
||||
with patch.object(
|
||||
persona_repository,
|
||||
"get_pool",
|
||||
side_effect=AssertionError("learner must be rejected before DB access"),
|
||||
):
|
||||
with self.assertRaises(ValueError):
|
||||
await persona_repository.list_persona_review_queue(role="learner")
|
||||
|
||||
async def test_learner_cannot_call_review_route(self) -> None:
|
||||
with patch.object(
|
||||
personas,
|
||||
"list_persona_review_queue",
|
||||
AsyncMock(side_effect=AssertionError("learner must not reach review repository")),
|
||||
) as review_queue:
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
await personas.list_persona_reviews(_principal(Role.LEARNER))
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 403)
|
||||
review_queue.assert_not_awaited()
|
||||
|
||||
async def test_teacher_review_route_returns_state_metadata(self) -> None:
|
||||
review_items = [
|
||||
PersonaReviewItem(
|
||||
persona_id="00000000-0000-0000-0000-000000000302",
|
||||
code="P2",
|
||||
version=2,
|
||||
status="draft",
|
||||
display_name="Draft Persona",
|
||||
difficulty="moderate",
|
||||
theory_target=["humanistic"],
|
||||
source_provenance="faculty import",
|
||||
is_synthetic=True,
|
||||
created_at="2026-01-01T00:00:00",
|
||||
approved_at=None,
|
||||
),
|
||||
PersonaReviewItem(
|
||||
persona_id="00000000-0000-0000-0000-000000000303",
|
||||
code="P3",
|
||||
version=1,
|
||||
status="review",
|
||||
display_name="Review Persona",
|
||||
difficulty="hard",
|
||||
theory_target=["cbt"],
|
||||
source_provenance="faculty import",
|
||||
is_synthetic=True,
|
||||
created_at="2026-01-02T00:00:00",
|
||||
approved_at=None,
|
||||
),
|
||||
]
|
||||
|
||||
with patch.object(
|
||||
personas,
|
||||
"list_persona_review_queue",
|
||||
AsyncMock(return_value=review_items),
|
||||
) as review_queue:
|
||||
response = await personas.list_persona_reviews(_principal(Role.TEACHER))
|
||||
|
||||
review_queue.assert_awaited_once_with(role="teacher")
|
||||
self.assertEqual([item.code for item in response], ["P2", "P3"])
|
||||
self.assertEqual([item.status for item in response], ["draft", "review"])
|
||||
self.assertEqual(response[0].version, 2)
|
||||
self.assertIsNone(response[0].approved_at)
|
||||
|
||||
async def test_admin_review_route_uses_admin_db_role(self) -> None:
|
||||
with patch.object(
|
||||
personas,
|
||||
"list_persona_review_queue",
|
||||
AsyncMock(return_value=[]),
|
||||
) as review_queue:
|
||||
response = await personas.list_persona_reviews(_principal(Role.ADMIN))
|
||||
|
||||
self.assertEqual(response, [])
|
||||
review_queue.assert_awaited_once_with(role="admin")
|
||||
|
||||
async def test_teacher_approves_review_persona_and_audits_decision(self) -> None:
|
||||
reviewer_id = "00000000-0000-0000-0000-000000000901"
|
||||
persona_id = "00000000-0000-0000-0000-000000000401"
|
||||
conn = _PersonaCardConn(
|
||||
[
|
||||
_card_row(
|
||||
persona_service.P2,
|
||||
persona_id=persona_id,
|
||||
status="review",
|
||||
version=2,
|
||||
),
|
||||
]
|
||||
)
|
||||
acquire_calls: list[dict[str, Any]] = []
|
||||
|
||||
def fake_acquire(**kwargs: Any) -> _Acquire:
|
||||
acquire_calls.append(kwargs)
|
||||
return _Acquire(conn)
|
||||
|
||||
with (
|
||||
patch.object(persona_repository, "get_pool", return_value=object()),
|
||||
patch.object(persona_repository, "acquire", fake_acquire),
|
||||
):
|
||||
updated = await persona_repository.update_persona_review_status(
|
||||
persona_id=persona_id,
|
||||
action="approve",
|
||||
reviewer_id=reviewer_id,
|
||||
role="teacher",
|
||||
)
|
||||
|
||||
self.assertIsNotNone(updated)
|
||||
assert updated is not None
|
||||
self.assertEqual(updated.status, "approved")
|
||||
self.assertEqual(updated.approved_at, "2026-01-03T00:00:00")
|
||||
self.assertEqual(acquire_calls, [{"role": "teacher", "user_id": reviewer_id}])
|
||||
update_query, update_args = conn.fetchrow_calls[0]
|
||||
self.assertIn("UPDATE app.persona_card", update_query)
|
||||
self.assertIn("status IN ('draft', 'review')", update_query)
|
||||
self.assertEqual(update_args, (persona_id, "approved", reviewer_id))
|
||||
audit_query, audit_args = conn.execute_calls[0]
|
||||
self.assertIn("INSERT INTO audit.audit_log", audit_query)
|
||||
self.assertEqual(audit_args[1], "persona_approve")
|
||||
self.assertEqual(audit_args[2], "persona_card")
|
||||
self.assertEqual(audit_args[3], persona_id)
|
||||
self.assertEqual(audit_args[4]["next_status"], "approved")
|
||||
|
||||
async def test_reject_review_persona_returns_it_to_draft_and_audits(self) -> None:
|
||||
reviewer_id = "00000000-0000-0000-0000-000000000901"
|
||||
persona_id = "00000000-0000-0000-0000-000000000402"
|
||||
conn = _PersonaCardConn(
|
||||
[
|
||||
_card_row(
|
||||
persona_service.P3,
|
||||
persona_id=persona_id,
|
||||
status="review",
|
||||
version=1,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(persona_repository, "get_pool", return_value=object()),
|
||||
patch.object(persona_repository, "acquire", lambda **_: _Acquire(conn)),
|
||||
):
|
||||
updated = await persona_repository.update_persona_review_status(
|
||||
persona_id=persona_id,
|
||||
action="reject",
|
||||
reviewer_id=reviewer_id,
|
||||
role="admin",
|
||||
)
|
||||
|
||||
self.assertIsNotNone(updated)
|
||||
assert updated is not None
|
||||
self.assertEqual(updated.status, "draft")
|
||||
self.assertIsNone(updated.approved_at)
|
||||
_, update_args = conn.fetchrow_calls[0]
|
||||
self.assertEqual(update_args, (persona_id, "draft", None))
|
||||
_, audit_args = conn.execute_calls[0]
|
||||
self.assertEqual(audit_args[1], "persona_reject")
|
||||
self.assertEqual(audit_args[4]["next_status"], "draft")
|
||||
|
||||
async def test_review_update_ignores_already_approved_persona(self) -> None:
|
||||
conn = _PersonaCardConn(
|
||||
[
|
||||
_card_row(
|
||||
persona_service.P1,
|
||||
persona_id="00000000-0000-0000-0000-000000000403",
|
||||
status="approved",
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(persona_repository, "get_pool", return_value=object()),
|
||||
patch.object(persona_repository, "acquire", lambda **_: _Acquire(conn)),
|
||||
):
|
||||
updated = await persona_repository.update_persona_review_status(
|
||||
persona_id="00000000-0000-0000-0000-000000000403",
|
||||
action="approve",
|
||||
reviewer_id="00000000-0000-0000-0000-000000000901",
|
||||
role="teacher",
|
||||
)
|
||||
|
||||
self.assertIsNone(updated)
|
||||
self.assertEqual(conn.execute_calls, [])
|
||||
|
||||
async def test_learner_cannot_call_review_decision_route(self) -> None:
|
||||
with patch.object(
|
||||
personas,
|
||||
"update_persona_review_status",
|
||||
AsyncMock(side_effect=AssertionError("learner must not reach review update")),
|
||||
) as update_review:
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
await personas.decide_persona_review(
|
||||
"00000000-0000-0000-0000-000000000404",
|
||||
personas.PersonaReviewDecisionRequest(action="approve"),
|
||||
_principal(Role.LEARNER),
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 403)
|
||||
update_review.assert_not_awaited()
|
||||
|
||||
async def test_review_decision_route_returns_404_for_non_pending_persona(self) -> None:
|
||||
with patch.object(
|
||||
personas,
|
||||
"update_persona_review_status",
|
||||
AsyncMock(return_value=None),
|
||||
) as update_review:
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
await personas.decide_persona_review(
|
||||
"00000000-0000-0000-0000-000000000405",
|
||||
personas.PersonaReviewDecisionRequest(action="approve"),
|
||||
_principal(Role.TEACHER),
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 404)
|
||||
update_review.assert_awaited_once_with(
|
||||
persona_id="00000000-0000-0000-0000-000000000405",
|
||||
action="approve",
|
||||
reviewer_id="00000000-0000-0000-0000-000000000901",
|
||||
role="teacher",
|
||||
)
|
||||
|
||||
async def test_review_decision_route_returns_updated_summary(self) -> None:
|
||||
updated_item = PersonaReviewItem(
|
||||
persona_id="00000000-0000-0000-0000-000000000406",
|
||||
code="P2",
|
||||
version=3,
|
||||
status="approved",
|
||||
display_name="Approved Persona",
|
||||
difficulty="moderate",
|
||||
theory_target=["humanistic"],
|
||||
source_provenance="faculty import",
|
||||
is_synthetic=True,
|
||||
created_at="2026-01-01T00:00:00",
|
||||
approved_at="2026-01-03T00:00:00",
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
personas,
|
||||
"update_persona_review_status",
|
||||
AsyncMock(return_value=updated_item),
|
||||
) as update_review:
|
||||
response = await personas.decide_persona_review(
|
||||
"00000000-0000-0000-0000-000000000406",
|
||||
personas.PersonaReviewDecisionRequest(action="approve"),
|
||||
_principal(Role.ADMIN),
|
||||
)
|
||||
|
||||
self.assertEqual(response.status, "approved")
|
||||
self.assertEqual(response.approved_at, "2026-01-03T00:00:00")
|
||||
update_review.assert_awaited_once_with(
|
||||
persona_id="00000000-0000-0000-0000-000000000406",
|
||||
action="approve",
|
||||
reviewer_id="00000000-0000-0000-0000-000000000901",
|
||||
role="admin",
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
440
apps/api/app/test_rbac_idor.py
Normal file
440
apps/api/app/test_rbac_idor.py
Normal file
|
|
@ -0,0 +1,440 @@
|
|||
"""Focused RBAC, IDOR, and session-read audit regression tests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from . import session_persistence
|
||||
from .deps import Principal, Role
|
||||
from .routes import sessions
|
||||
from .services import persona as persona_service, state_machine
|
||||
from .store import InProcSession, TurnRecord, store
|
||||
|
||||
|
||||
def _principal(
|
||||
*,
|
||||
user_id: str,
|
||||
role: Role = Role.LEARNER,
|
||||
) -> Principal:
|
||||
return Principal(
|
||||
user_id=user_id,
|
||||
role=role,
|
||||
cohort_ids=["cohort-a"] if role == Role.TEACHER else [],
|
||||
email=f"{role.value}-{user_id[-4:]}@example.test",
|
||||
display_name=f"{role.value.title()} {user_id[-4:]}",
|
||||
)
|
||||
|
||||
|
||||
def _session(
|
||||
*,
|
||||
session_id: str,
|
||||
learner_id: str,
|
||||
) -> InProcSession:
|
||||
card = persona_service.P1
|
||||
return InProcSession(
|
||||
session_id=session_id,
|
||||
case_id=f"case-{session_id[-12:]}",
|
||||
learner_id=learner_id,
|
||||
persona_code=card.code,
|
||||
theory_mode="humanistic",
|
||||
persona=card,
|
||||
state=state_machine.SessionState(
|
||||
resistance=card.base_resistance(),
|
||||
ideation_stage=card.ideation_baseline(),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _turn(
|
||||
*,
|
||||
seq: int,
|
||||
speaker: str,
|
||||
text: str,
|
||||
visible_to: tuple[str, ...] = ("client", "counselor", "evaluator"),
|
||||
) -> TurnRecord:
|
||||
return TurnRecord(
|
||||
turn_seq=seq,
|
||||
speaker=speaker,
|
||||
stage="rapport",
|
||||
text=text,
|
||||
text_masked=text,
|
||||
created_at=1_800_000_000.0 + seq,
|
||||
visible_to=visible_to,
|
||||
)
|
||||
|
||||
|
||||
class _Acquire:
|
||||
def __init__(self, conn: "_FakeConn") -> None:
|
||||
self.conn = conn
|
||||
|
||||
async def __aenter__(self) -> "_FakeConn":
|
||||
return self.conn
|
||||
|
||||
async def __aexit__(self, exc_type: object, exc: object, tb: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class _FakeConn:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
fetchrow_results: list[Any] | None = None,
|
||||
fetch_results: list[list[Any]] | None = None,
|
||||
) -> None:
|
||||
self.fetchrow_results = list(fetchrow_results or [])
|
||||
self.fetch_results = list(fetch_results or [])
|
||||
self.executed: list[tuple[str, tuple[Any, ...]]] = []
|
||||
|
||||
async def fetchrow(self, *args: Any, **kwargs: Any) -> Any:
|
||||
if not self.fetchrow_results:
|
||||
raise AssertionError("unexpected fetchrow")
|
||||
return self.fetchrow_results.pop(0)
|
||||
|
||||
async def fetch(self, *args: Any, **kwargs: Any) -> list[Any]:
|
||||
if not self.fetch_results:
|
||||
raise AssertionError("unexpected fetch")
|
||||
return self.fetch_results.pop(0)
|
||||
|
||||
async def execute(self, query: str, *args: Any) -> str:
|
||||
self.executed.append((query, args))
|
||||
return "INSERT 0 1"
|
||||
|
||||
|
||||
def _audit_calls(conn: _FakeConn) -> list[tuple[str, tuple[Any, ...]]]:
|
||||
return [
|
||||
call
|
||||
for call in conn.executed
|
||||
if "INSERT INTO audit.audit_log" in call[0]
|
||||
]
|
||||
|
||||
|
||||
class LearnerSessionIdorTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self) -> None:
|
||||
store._sessions.clear()
|
||||
|
||||
async def asyncTearDown(self) -> None:
|
||||
store._sessions.clear()
|
||||
|
||||
async def test_get_session_detail_rejects_other_learner_session_id(self) -> None:
|
||||
owner = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000101",
|
||||
)
|
||||
intruder = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000202",
|
||||
)
|
||||
sess = _session(
|
||||
session_id="00000000-0000-0000-0000-00000000a101",
|
||||
learner_id=owner.user_id,
|
||||
)
|
||||
sess.turns.append(
|
||||
_turn(
|
||||
seq=1,
|
||||
speaker="counselor",
|
||||
text="owner-visible turn must not grant access",
|
||||
visible_to=("counselor", "evaluator"),
|
||||
)
|
||||
)
|
||||
store.put(sess)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
sessions.session_persistence,
|
||||
"load_session",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch.object(sessions, "runtime_fallback_allowed", return_value=True),
|
||||
patch.object(
|
||||
sessions,
|
||||
"_review_ready",
|
||||
AsyncMock(side_effect=AssertionError("review lookup should not run")),
|
||||
) as review_ready,
|
||||
):
|
||||
with self.assertRaises(HTTPException) as caught:
|
||||
await sessions.get_session_detail(sess.session_id, intruder)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 403)
|
||||
self.assertIn("does not belong", caught.exception.detail)
|
||||
review_ready.assert_not_awaited()
|
||||
|
||||
async def test_session_detail_filters_evaluator_only_turns(self) -> None:
|
||||
owner = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000111",
|
||||
)
|
||||
sess = _session(
|
||||
session_id="00000000-0000-0000-0000-00000000e111",
|
||||
learner_id=owner.user_id,
|
||||
)
|
||||
sess.turns.extend(
|
||||
[
|
||||
_turn(seq=1, speaker="counselor", text="learner normal turn"),
|
||||
_turn(seq=2, speaker="client", text="client normal turn"),
|
||||
_turn(
|
||||
seq=3,
|
||||
speaker="counselor",
|
||||
text="SECRET_EVALUATOR_ONLY_DETAIL",
|
||||
visible_to=("evaluator",),
|
||||
),
|
||||
]
|
||||
)
|
||||
store.put(sess)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
sessions.session_persistence,
|
||||
"load_session",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch.object(sessions, "runtime_fallback_allowed", return_value=True),
|
||||
patch.object(sessions, "_review_ready", AsyncMock(return_value=False)),
|
||||
):
|
||||
response = await sessions.get_session_detail(sess.session_id, owner)
|
||||
|
||||
self.assertEqual([turn.text for turn in response.turns], ["learner normal turn", "client normal turn"])
|
||||
self.assertEqual([turn.speaker for turn in response.turns], ["learner", "client"])
|
||||
|
||||
async def test_session_review_filters_evaluator_only_turns_and_payload(self) -> None:
|
||||
owner = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000112",
|
||||
)
|
||||
sess = _session(
|
||||
session_id="00000000-0000-0000-0000-00000000e112",
|
||||
learner_id=owner.user_id,
|
||||
)
|
||||
sess.ended = True
|
||||
sess.ended_at = 1_800_000_120.0
|
||||
sess.turns.extend(
|
||||
[
|
||||
_turn(seq=1, speaker="counselor", text="visible learner review turn"),
|
||||
_turn(seq=2, speaker="client", text="visible client review turn"),
|
||||
_turn(
|
||||
seq=3,
|
||||
speaker="client",
|
||||
text="SECRET_EVALUATOR_ONLY_REVIEW",
|
||||
visible_to=("evaluator",),
|
||||
),
|
||||
]
|
||||
)
|
||||
store.put(sess)
|
||||
evaluation_record = {
|
||||
"status": "ready",
|
||||
"payload": {
|
||||
"strengths": ["SECRET_EVALUATOR_ONLY_REVIEW"],
|
||||
"improvements": ["SECRET_EVALUATOR_ONLY_REVIEW"],
|
||||
"supervisor_rationale": "SECRET_EVALUATOR_ONLY_REVIEW",
|
||||
"alternative_utterances": ["SECRET_EVALUATOR_ONLY_REVIEW"],
|
||||
},
|
||||
}
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
sessions.session_persistence,
|
||||
"load_session",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch.object(sessions, "runtime_fallback_allowed", return_value=True),
|
||||
patch.object(
|
||||
sessions.session_persistence,
|
||||
"load_session_evaluation",
|
||||
AsyncMock(return_value=(evaluation_record, True)),
|
||||
),
|
||||
):
|
||||
response = await sessions.get_session_review(sess.session_id, owner)
|
||||
|
||||
self.assertEqual([turn.text for turn in response.turns], ["visible learner review turn", "visible client review turn"])
|
||||
self.assertFalse(response.reviewReady)
|
||||
rendered = " ".join(
|
||||
[
|
||||
response.summary,
|
||||
response.clientFeedback or "",
|
||||
response.nextLine or "",
|
||||
*[turn.text for turn in response.turns],
|
||||
*[point.body for point in response.goodMoments],
|
||||
*[point.body for point in response.growthPoints],
|
||||
]
|
||||
)
|
||||
self.assertNotIn("SECRET_EVALUATOR_ONLY_REVIEW", rendered)
|
||||
|
||||
async def test_submit_turn_sends_only_client_visible_history_to_engine(self) -> None:
|
||||
owner = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000113",
|
||||
)
|
||||
sess = _session(
|
||||
session_id="00000000-0000-0000-0000-00000000e113",
|
||||
learner_id=owner.user_id,
|
||||
)
|
||||
sess.turns.extend(
|
||||
[
|
||||
_turn(seq=1, speaker="counselor", text="client-visible history"),
|
||||
_turn(
|
||||
seq=2,
|
||||
speaker="client",
|
||||
text="SECRET_EVALUATOR_ONLY_ENGINE_CONTEXT",
|
||||
visible_to=("evaluator",),
|
||||
),
|
||||
]
|
||||
)
|
||||
store.put(sess)
|
||||
captured_recent_turns: list[dict[str, str]] | None = None
|
||||
|
||||
async def successful_turn(ctx, engine, **kwargs):
|
||||
nonlocal captured_recent_turns
|
||||
captured_recent_turns = list(ctx.recent_turns)
|
||||
assert ctx.state_after is not None
|
||||
return sessions.orchestrator.TurnResult(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
stage=ctx.state_after.stage.value,
|
||||
effective_openness=ctx.state_after.effective_openness,
|
||||
client_reply="client reply",
|
||||
safety_flagged=False,
|
||||
state_after=ctx.state_after,
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
sessions.session_persistence,
|
||||
"load_session",
|
||||
AsyncMock(return_value=None),
|
||||
),
|
||||
patch.object(sessions, "runtime_fallback_allowed", return_value=True),
|
||||
patch.object(sessions.orchestrator, "run_turn_generate", successful_turn),
|
||||
):
|
||||
await sessions.submit_turn(
|
||||
sess.session_id,
|
||||
sessions.TurnRequest(text="new learner turn"),
|
||||
owner,
|
||||
)
|
||||
|
||||
self.assertEqual(captured_recent_turns, [{"speaker": "counselor", "text": "client-visible history"}])
|
||||
self.assertNotIn(
|
||||
"SECRET_EVALUATOR_ONLY_ENGINE_CONTEXT",
|
||||
" ".join(turn["text"] for turn in captured_recent_turns or []),
|
||||
)
|
||||
|
||||
async def test_learner_session_list_filters_other_runtime_sessions(self) -> None:
|
||||
owner = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000303",
|
||||
)
|
||||
other = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000404",
|
||||
)
|
||||
owned_session = _session(
|
||||
session_id="00000000-0000-0000-0000-00000000b303",
|
||||
learner_id=owner.user_id,
|
||||
)
|
||||
other_session = _session(
|
||||
session_id="00000000-0000-0000-0000-00000000b404",
|
||||
learner_id=other.user_id,
|
||||
)
|
||||
store.put(owned_session)
|
||||
store.put(other_session)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
sessions.session_persistence,
|
||||
"list_sessions",
|
||||
AsyncMock(return_value=([], False)),
|
||||
),
|
||||
patch.object(sessions, "require_runtime_fallback_allowed", return_value=None),
|
||||
patch.object(sessions, "_review_ready", AsyncMock(return_value=False)),
|
||||
):
|
||||
response = await sessions.list_learner_sessions(owner)
|
||||
|
||||
self.assertEqual(response.source, "runtime")
|
||||
self.assertEqual([item.session_id for item in response.sessions], [owned_session.session_id])
|
||||
|
||||
|
||||
class TeacherAdminAuditTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_teacher_list_sessions_inserts_read_audit_log(self) -> None:
|
||||
teacher = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000505",
|
||||
role=Role.TEACHER,
|
||||
)
|
||||
session_id = "00000000-0000-0000-0000-00000000c505"
|
||||
row = {"id": session_id}
|
||||
sess = _session(
|
||||
session_id=session_id,
|
||||
learner_id="00000000-0000-0000-0000-000000000606",
|
||||
)
|
||||
conn = _FakeConn(
|
||||
fetchrow_results=[None],
|
||||
fetch_results=[[row], []],
|
||||
)
|
||||
acquire_calls: list[dict[str, Any]] = []
|
||||
|
||||
def fake_acquire(**kwargs: Any) -> _Acquire:
|
||||
acquire_calls.append(kwargs)
|
||||
return _Acquire(conn)
|
||||
|
||||
with (
|
||||
patch.object(session_persistence, "get_pool", return_value=object()),
|
||||
patch.object(session_persistence, "acquire", fake_acquire),
|
||||
patch.object(session_persistence, "_session_from_rows", return_value=sess),
|
||||
):
|
||||
found, durable = await session_persistence.list_sessions(teacher)
|
||||
|
||||
self.assertTrue(durable)
|
||||
self.assertEqual(found, [sess])
|
||||
self.assertEqual(acquire_calls[0]["role"], "teacher")
|
||||
audit = _audit_calls(conn)
|
||||
self.assertEqual(len(audit), 1)
|
||||
_, args = audit[0]
|
||||
self.assertEqual(args[0], teacher.user_id)
|
||||
self.assertEqual(args[1], "read_session")
|
||||
self.assertEqual(args[2], "session_list")
|
||||
self.assertEqual(args[3], "sessions")
|
||||
self.assertEqual(args[4]["access"], "list_sessions")
|
||||
self.assertEqual(args[4]["role"], "teacher")
|
||||
self.assertEqual(args[4]["result_count"], 1)
|
||||
|
||||
async def test_admin_load_session_inserts_read_audit_log(self) -> None:
|
||||
admin = _principal(
|
||||
user_id="00000000-0000-0000-0000-000000000707",
|
||||
role=Role.ADMIN,
|
||||
)
|
||||
session_id = "00000000-0000-0000-0000-00000000d707"
|
||||
learner_id = "00000000-0000-0000-0000-000000000808"
|
||||
sess = _session(session_id=session_id, learner_id=learner_id)
|
||||
conn = _FakeConn(
|
||||
fetchrow_results=[
|
||||
{"id": session_id, "ended_at": None},
|
||||
None,
|
||||
],
|
||||
fetch_results=[[]],
|
||||
)
|
||||
acquire_calls: list[dict[str, Any]] = []
|
||||
|
||||
def fake_acquire(**kwargs: Any) -> _Acquire:
|
||||
acquire_calls.append(kwargs)
|
||||
return _Acquire(conn)
|
||||
|
||||
with (
|
||||
patch.object(session_persistence, "get_pool", return_value=object()),
|
||||
patch.object(session_persistence, "acquire", fake_acquire),
|
||||
patch.object(session_persistence, "_session_from_rows", return_value=sess),
|
||||
):
|
||||
found = await session_persistence.load_session(
|
||||
session_id,
|
||||
admin,
|
||||
allow_ended=True,
|
||||
)
|
||||
|
||||
self.assertEqual(found, sess)
|
||||
self.assertEqual(acquire_calls[0]["role"], "admin")
|
||||
audit = _audit_calls(conn)
|
||||
self.assertEqual(len(audit), 1)
|
||||
_, args = audit[0]
|
||||
self.assertEqual(args[0], admin.user_id)
|
||||
self.assertEqual(args[1], "read_session")
|
||||
self.assertEqual(args[2], "session")
|
||||
self.assertEqual(args[3], session_id)
|
||||
self.assertEqual(args[4]["access"], "load_session")
|
||||
self.assertEqual(args[4]["role"], "admin")
|
||||
self.assertEqual(args[4]["learner_id"], learner_id)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
@ -264,6 +264,42 @@ class RuntimeFallbackPolicyTest(unittest.IsolatedAsyncioTestCase):
|
|||
self.assertEqual(cfg.environment, "staging")
|
||||
self.assertEqual(cfg.cors_origins, ["https://vignette.chanpaca.net"])
|
||||
|
||||
def test_non_dev_accepts_explicit_local_vite_cors_ports(self) -> None:
|
||||
cfg = Settings(
|
||||
environment="prod",
|
||||
auth_dev_login_enabled=False,
|
||||
auto_seed_personas=False,
|
||||
allow_seed_persona_fallback=False,
|
||||
session_secret="prod-secret-change-me",
|
||||
oauth_google_client_id="google-client-id",
|
||||
oauth_google_client_secret="google-client-secret",
|
||||
frontend_base_url="https://vignette.chanpaca.net",
|
||||
cors_origins=[
|
||||
"https://vignette.chanpaca.net",
|
||||
"http://localhost:5170",
|
||||
"http://127.0.0.1:5180",
|
||||
],
|
||||
)
|
||||
|
||||
self.assertIn("http://localhost:5170", cfg.cors_origins)
|
||||
self.assertIn("http://127.0.0.1:5180", cfg.cors_origins)
|
||||
|
||||
def test_non_dev_rejects_unscoped_local_cors_ports(self) -> None:
|
||||
with self.assertRaises(ValueError) as caught:
|
||||
Settings(
|
||||
environment="prod",
|
||||
auth_dev_login_enabled=False,
|
||||
auto_seed_personas=False,
|
||||
allow_seed_persona_fallback=False,
|
||||
session_secret="prod-secret-change-me",
|
||||
oauth_google_client_id="google-client-id",
|
||||
oauth_google_client_secret="google-client-secret",
|
||||
frontend_base_url="https://vignette.chanpaca.net",
|
||||
cors_origins=["http://localhost:5181"],
|
||||
)
|
||||
|
||||
self.assertIn("CORS_ORIGINS", str(caught.exception))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from .engine_client import EngineError
|
|||
from .routes import sessions
|
||||
from .routes import voice as voice_routes
|
||||
from .services import orchestrator, persona as persona_service, state_machine
|
||||
from .services.voice import VoicePreset
|
||||
from .services.voice import TTSChunk, TranscriptResult, VoicePreset
|
||||
from .store import InProcSession, store
|
||||
|
||||
|
||||
|
|
@ -81,6 +81,152 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
|||
self.assertEqual(caught.exception.status_code, 503)
|
||||
self.assertEqual(sess.turns, [])
|
||||
|
||||
async def test_generate_turn_persists_client_engine_telemetry(self) -> None:
|
||||
principal = _principal()
|
||||
sess = _session(principal)
|
||||
|
||||
async def successful_turn(ctx, engine, **kwargs):
|
||||
assert ctx.state_after is not None
|
||||
return orchestrator.TurnResult(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
stage=ctx.state_after.stage.value,
|
||||
effective_openness=ctx.state_after.effective_openness,
|
||||
client_reply="괜찮아요. 천천히 말해볼게요.",
|
||||
safety_flagged=False,
|
||||
state_after=ctx.state_after,
|
||||
llm_provider="claude_cli",
|
||||
model="gateway-default",
|
||||
tokens_in=17,
|
||||
tokens_out=23,
|
||||
cost_usd=0.012345,
|
||||
)
|
||||
|
||||
with patch.object(sessions.orchestrator, "run_turn_generate", successful_turn):
|
||||
response = await sessions.submit_turn(
|
||||
sess.session_id,
|
||||
sessions.TurnRequest(text="요즘 많이 힘들었겠어요."),
|
||||
principal,
|
||||
)
|
||||
|
||||
self.assertEqual(response.client_reply, "괜찮아요. 천천히 말해볼게요.")
|
||||
self.assertEqual(len(sess.turns), 2)
|
||||
learner_turn, client_turn = sess.turns
|
||||
self.assertIsNone(learner_turn.llm_provider)
|
||||
self.assertEqual(client_turn.llm_provider, "claude_cli")
|
||||
self.assertEqual(client_turn.model, "gateway-default")
|
||||
self.assertEqual(client_turn.tokens_in, 17)
|
||||
self.assertEqual(client_turn.tokens_out, 23)
|
||||
self.assertEqual(client_turn.cost_usd, 0.012345)
|
||||
|
||||
async def test_stream_turn_persists_client_engine_telemetry(self) -> None:
|
||||
principal = _principal()
|
||||
sess = _session(principal)
|
||||
|
||||
async def successful_stream(ctx, engine):
|
||||
assert ctx.state_after is not None
|
||||
yield orchestrator.StreamEvent("token", {"text": "괜찮아요."})
|
||||
yield orchestrator.StreamEvent(
|
||||
"done",
|
||||
{
|
||||
"session_id": ctx.session_id,
|
||||
"stage": ctx.state_after.stage.value,
|
||||
"effective_openness": ctx.state_after.effective_openness,
|
||||
"turn_seq": ctx.state_after.turn_seq,
|
||||
"safety_flagged": False,
|
||||
"llm_provider": "claude_cli",
|
||||
"model": "gateway-default",
|
||||
"tokens_in": 31,
|
||||
"tokens_out": 37,
|
||||
"cost_usd": 0.023456,
|
||||
},
|
||||
)
|
||||
|
||||
with patch.object(sessions.orchestrator, "run_turn_stream", successful_stream):
|
||||
response = await sessions.stream_turn(
|
||||
sess.session_id,
|
||||
sessions.TurnRequest(text="스트림 성공 발화"),
|
||||
principal,
|
||||
)
|
||||
body = await _consume_event_source(response)
|
||||
|
||||
self.assertIn(b"done", body)
|
||||
self.assertEqual(len(sess.turns), 2)
|
||||
client_turn = sess.turns[1]
|
||||
self.assertEqual(client_turn.llm_provider, "claude_cli")
|
||||
self.assertEqual(client_turn.model, "gateway-default")
|
||||
self.assertEqual(client_turn.tokens_in, 31)
|
||||
self.assertEqual(client_turn.tokens_out, 37)
|
||||
self.assertEqual(client_turn.cost_usd, 0.023456)
|
||||
|
||||
async def test_run_turn_stream_parses_gateway_done_telemetry(self) -> None:
|
||||
class FakeStreamEngine:
|
||||
engine_mode = "claude_cli"
|
||||
default_model = None
|
||||
|
||||
async def stream(self, req):
|
||||
yield "event: token"
|
||||
yield '{"ignored":"not data"}'
|
||||
yield 'data: {"text":"부분 응답"}'
|
||||
yield "event: done"
|
||||
yield (
|
||||
'data: {"provider":"claude_cli","model":"gateway-default",'
|
||||
'"tokens_in":5,"tokens_out":7,"cost_usd":0.034567}'
|
||||
)
|
||||
|
||||
principal = _principal()
|
||||
sess = _session(principal)
|
||||
ctx = orchestrator.prepare_turn(
|
||||
session_id=sess.session_id,
|
||||
case_id=sess.case_id,
|
||||
card=sess.persona,
|
||||
state=sess.state,
|
||||
learner_text="게이트웨이 스트림 테스트",
|
||||
recent_turns=[],
|
||||
)
|
||||
|
||||
events = [
|
||||
event
|
||||
async for event in orchestrator.run_turn_stream(ctx, FakeStreamEngine()) # type: ignore[arg-type]
|
||||
]
|
||||
|
||||
self.assertEqual([event.event for event in events], ["token", "done"])
|
||||
self.assertEqual(events[0].data["text"], "부분 응답")
|
||||
self.assertEqual(events[1].data["llm_provider"], "claude_cli")
|
||||
self.assertEqual(events[1].data["model"], "gateway-default")
|
||||
self.assertEqual(events[1].data["tokens_in"], 5)
|
||||
self.assertEqual(events[1].data["tokens_out"], 7)
|
||||
self.assertEqual(events[1].data["cost_usd"], 0.034567)
|
||||
|
||||
async def test_run_turn_stream_treats_gateway_error_event_as_error(self) -> None:
|
||||
class FakeStreamEngine:
|
||||
engine_mode = "claude_cli"
|
||||
default_model = None
|
||||
|
||||
async def stream(self, req):
|
||||
yield "event: token"
|
||||
yield 'data: {"text":"부분 응답"}'
|
||||
yield "event: error"
|
||||
yield 'data: {"detail":"engine unavailable: gateway"}'
|
||||
|
||||
principal = _principal()
|
||||
sess = _session(principal)
|
||||
ctx = orchestrator.prepare_turn(
|
||||
session_id=sess.session_id,
|
||||
case_id=sess.case_id,
|
||||
card=sess.persona,
|
||||
state=sess.state,
|
||||
learner_text="게이트웨이 오류 테스트",
|
||||
recent_turns=[],
|
||||
)
|
||||
|
||||
events = [
|
||||
event
|
||||
async for event in orchestrator.run_turn_stream(ctx, FakeStreamEngine()) # type: ignore[arg-type]
|
||||
]
|
||||
|
||||
self.assertEqual([event.event for event in events], ["token", "error"])
|
||||
self.assertIn("engine unavailable", events[1].data["detail"])
|
||||
|
||||
async def test_stream_turn_engine_error_event_does_not_append_partial_turns(self) -> None:
|
||||
principal = _principal()
|
||||
sess = _session(principal)
|
||||
|
|
@ -138,6 +284,79 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
|||
)
|
||||
self.assertEqual(sess.turns, [])
|
||||
|
||||
async def test_voice_audio_turn_persists_paralinguistic_metadata(self) -> None:
|
||||
class FakeWebSocket:
|
||||
def __init__(self) -> None:
|
||||
self.messages: list[dict[str, object]] = []
|
||||
self.binary: list[bytes] = []
|
||||
self.client_state = voice_routes.WebSocketState.CONNECTED
|
||||
|
||||
async def send_text(self, data: str) -> None:
|
||||
import json
|
||||
|
||||
self.messages.append(json.loads(data))
|
||||
|
||||
async def send_bytes(self, data: bytes) -> None:
|
||||
self.binary.append(data)
|
||||
|
||||
async def successful_turn(ctx, engine, **kwargs):
|
||||
assert ctx.state_after is not None
|
||||
return orchestrator.TurnResult(
|
||||
turn_seq=ctx.state_after.turn_seq,
|
||||
stage=ctx.state_after.stage.value,
|
||||
effective_openness=ctx.state_after.effective_openness,
|
||||
client_reply="천천히 말해줘서 고마워요.",
|
||||
safety_flagged=False,
|
||||
state_after=ctx.state_after,
|
||||
llm_provider="claude_cli",
|
||||
model="gateway-default",
|
||||
tokens_in=11,
|
||||
tokens_out=13,
|
||||
cost_usd=0.0012,
|
||||
)
|
||||
|
||||
async def fake_synthesize_stream(text, voice_preset):
|
||||
yield TTSChunk(audio=b"tts-audio")
|
||||
|
||||
principal = _principal()
|
||||
sess = _session(principal)
|
||||
websocket = FakeWebSocket()
|
||||
audio = b"\x00\x80" * 1600
|
||||
|
||||
with patch.object(
|
||||
voice_routes.voice_service,
|
||||
"transcribe",
|
||||
AsyncMock(return_value=TranscriptResult(text="오늘은 좀 힘들었어요.", duration=2.0)),
|
||||
), patch.object(
|
||||
voice_routes.orchestrator,
|
||||
"run_turn_generate",
|
||||
successful_turn,
|
||||
), patch.object(
|
||||
voice_routes.voice_service,
|
||||
"synthesize_stream",
|
||||
fake_synthesize_stream,
|
||||
):
|
||||
await voice_routes._handle_utterance(
|
||||
websocket, # type: ignore[arg-type]
|
||||
session_id=sess.session_id,
|
||||
principal=principal,
|
||||
voice_preset=VoicePreset(preset="neutral", openai_voice="sage"),
|
||||
audio=audio,
|
||||
fmt="webm",
|
||||
silence_ms=1234,
|
||||
barge_in=True,
|
||||
)
|
||||
|
||||
self.assertEqual(len(sess.turns), 2)
|
||||
learner_turn, client_turn = sess.turns
|
||||
self.assertTrue(str(learner_turn.audio_ref).startswith("voice:webm:sha256:"))
|
||||
self.assertEqual(learner_turn.silence_ms, 1234)
|
||||
self.assertGreater(learner_turn.speech_rate or 0, 0)
|
||||
self.assertTrue(learner_turn.barge_in)
|
||||
self.assertIsNone(client_turn.audio_ref)
|
||||
self.assertEqual(client_turn.llm_provider, "claude_cli")
|
||||
self.assertTrue(any(message.get("type") == "tts_end" for message in websocket.messages))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
|
|
|||
73
apps/api/app/test_state_machine_resistance.py
Normal file
73
apps/api/app/test_state_machine_resistance.py
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
"""Regression tests for the deterministic resistance engine."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from .services import state_machine
|
||||
from .services.persona import P1
|
||||
|
||||
|
||||
EMPATHIC_UTTERANCES = [
|
||||
"얼마나 힘들었는지 마음이 느껴져요. 어떤 순간이 제일 버거웠나요?",
|
||||
"그런 마음을 꺼내는 것 자체가 쉽지 않았을 것 같아요. 더 말해줘도 괜찮아요.",
|
||||
"잠도 잘 못 자고 학교도 버거웠다면 하루가 길게 느껴졌겠어요.",
|
||||
"지금은 해결책보다 그 마음을 천천히 이해하는 게 먼저인 것 같아요.",
|
||||
"그 시간을 버텨온 마음을 함께 살펴보고 싶어요. 무엇부터 이야기해볼까요?",
|
||||
]
|
||||
|
||||
ADVICE_JUMP_UTTERANCES = [
|
||||
"그냥 학교는 가야 해요. 노력하면 하면 돼요. 왜 안 하죠?",
|
||||
"그건 잘못 생각하는 거예요. 원래 다 힘들어요.",
|
||||
"당연히 엄마 말을 들어야죠. 하지 마세요.",
|
||||
"내 생각엔 그냥 계획표를 만들면 돼요.",
|
||||
"그러니까 더 노력해야 해요. 왜 안 바꾸나요?",
|
||||
]
|
||||
|
||||
|
||||
def _initial_p1_state() -> state_machine.SessionState:
|
||||
return state_machine.init_state(
|
||||
params=P1.openness_params(),
|
||||
)
|
||||
|
||||
|
||||
def _run_curve(utterances: list[str]) -> list[state_machine.SessionState]:
|
||||
state = _initial_p1_state()
|
||||
curve: list[state_machine.SessionState] = []
|
||||
for utterance in utterances:
|
||||
signal = state_machine.estimate_rapport_signal(utterance)
|
||||
state = state_machine.evolve(
|
||||
state,
|
||||
rapport_signal=signal,
|
||||
unlock_rate=P1.unlock_rate(),
|
||||
decay_floor=P1.decay_floor(),
|
||||
)
|
||||
curve.append(state)
|
||||
return curve
|
||||
|
||||
|
||||
class ResistanceEngineTest(unittest.TestCase):
|
||||
def test_empathy_opens_p1_while_advice_jump_closes_it(self) -> None:
|
||||
empathy_curve = _run_curve(EMPATHIC_UTTERANCES)
|
||||
advice_curve = _run_curve(ADVICE_JUMP_UTTERANCES)
|
||||
empathy_final = empathy_curve[-1]
|
||||
advice_final = advice_curve[-1]
|
||||
|
||||
self.assertGreater(empathy_final.rapport_credit, advice_final.rapport_credit)
|
||||
self.assertLess(empathy_final.resistance, advice_final.resistance)
|
||||
self.assertGreater(empathy_final.effective_openness, advice_final.effective_openness)
|
||||
self.assertEqual(empathy_final.stage, state_machine.Stage.EXPLORE)
|
||||
self.assertEqual(advice_final.stage, state_machine.Stage.RAPPORT)
|
||||
self.assertGreater(empathy_final.effective_openness, 0.1)
|
||||
self.assertEqual(advice_final.effective_openness, 0.0)
|
||||
|
||||
def test_advice_jump_never_advances_stage_after_five_turns(self) -> None:
|
||||
advice_curve = _run_curve(ADVICE_JUMP_UTTERANCES)
|
||||
|
||||
self.assertTrue(all(state.stage is state_machine.Stage.RAPPORT for state in advice_curve))
|
||||
self.assertTrue(all(state.rapport_credit == 0 for state in advice_curve))
|
||||
self.assertGreaterEqual(advice_curve[-1].resistance, 0.95)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
223
apps/api/app/test_voice_service.py
Normal file
223
apps/api/app/test_voice_service.py
Normal file
|
|
@ -0,0 +1,223 @@
|
|||
"""Deterministic tests for voice preset, TTS payload, and EOT helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import unittest
|
||||
|
||||
from .services.voice import (
|
||||
DEFAULT_OPENAI_VOICE,
|
||||
EOT_SILENCE_THRESHOLD_MS,
|
||||
TTS_ENDPOINT,
|
||||
TTS_MODEL,
|
||||
TTS_MODEL_FALLBACK,
|
||||
VoicePreset,
|
||||
VoiceService,
|
||||
assess_end_of_turn,
|
||||
build_tts_payload,
|
||||
resolve_voice,
|
||||
)
|
||||
|
||||
|
||||
class VoicePresetResolutionTest(unittest.TestCase):
|
||||
def test_persona_codes_resolve_distinct_voice_presets(self) -> None:
|
||||
expected = {
|
||||
"P1": ("soft-young-fem", "coral", 0.96),
|
||||
"P2": ("calm-adult-male", "ash", 1.0),
|
||||
"P3": ("warm-adult-fem", "shimmer", 0.98),
|
||||
}
|
||||
|
||||
resolved = {code: resolve_voice(persona_code=code) for code in expected}
|
||||
|
||||
self.assertEqual(
|
||||
{voice.openai_voice for voice in resolved.values()},
|
||||
{"coral", "ash", "shimmer"},
|
||||
)
|
||||
for code, (preset, openai_voice, rate) in expected.items():
|
||||
with self.subTest(code=code):
|
||||
voice = resolved[code]
|
||||
self.assertEqual(voice.preset, preset)
|
||||
self.assertEqual(voice.openai_voice, openai_voice)
|
||||
self.assertAlmostEqual(voice.rate, rate)
|
||||
|
||||
self.assertEqual(resolve_voice(persona_code="p2").preset, "calm-adult-male")
|
||||
|
||||
def test_invalid_explicit_preset_falls_back_to_default_openai_voice(self) -> None:
|
||||
voice = resolve_voice(persona_code="P1", preset="not-a-real-preset")
|
||||
|
||||
self.assertEqual(voice.preset, "not-a-real-preset")
|
||||
self.assertEqual(voice.openai_voice, DEFAULT_OPENAI_VOICE)
|
||||
self.assertEqual(voice.rate, 1.0)
|
||||
|
||||
|
||||
class TTSPayloadTest(unittest.TestCase):
|
||||
def test_payload_contains_openai_tts_fields_and_clamps_high_speed(self) -> None:
|
||||
voice = VoicePreset(
|
||||
preset="soft-young-fem",
|
||||
openai_voice="coral",
|
||||
rate=9.5,
|
||||
instructions="Speak gently with low intensity.",
|
||||
)
|
||||
|
||||
payload = build_tts_payload("Client reply", voice)
|
||||
|
||||
self.assertEqual(
|
||||
payload,
|
||||
{
|
||||
"model": TTS_MODEL,
|
||||
"voice": "coral",
|
||||
"input": "Client reply",
|
||||
"response_format": "mp3",
|
||||
"speed": 4.0,
|
||||
"instructions": "Speak gently with low intensity.",
|
||||
},
|
||||
)
|
||||
|
||||
def test_payload_omits_instructions_for_fallback_model_and_clamps_low_speed(self) -> None:
|
||||
voice = VoicePreset(
|
||||
preset="neutral",
|
||||
openai_voice="sage",
|
||||
rate=0.1,
|
||||
instructions="This should not be sent to tts-1.",
|
||||
)
|
||||
|
||||
payload = build_tts_payload(
|
||||
"Fallback reply",
|
||||
voice,
|
||||
model=TTS_MODEL_FALLBACK,
|
||||
response_format="opus",
|
||||
)
|
||||
|
||||
self.assertEqual(payload["model"], TTS_MODEL_FALLBACK)
|
||||
self.assertEqual(payload["voice"], "sage")
|
||||
self.assertEqual(payload["input"], "Fallback reply")
|
||||
self.assertEqual(payload["response_format"], "opus")
|
||||
self.assertEqual(payload["speed"], 0.25)
|
||||
self.assertNotIn("instructions", payload)
|
||||
|
||||
def test_payload_uses_default_speed_for_invalid_rate(self) -> None:
|
||||
voice = VoicePreset(
|
||||
preset="neutral",
|
||||
openai_voice="sage",
|
||||
rate="fast", # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
payload = build_tts_payload("Client reply", voice)
|
||||
|
||||
self.assertEqual(payload["speed"], 1.0)
|
||||
|
||||
|
||||
class _FakeTTSStream:
|
||||
def __init__(self, chunks: list[bytes]) -> None:
|
||||
self.status_code = 200
|
||||
self._chunks = chunks
|
||||
self.closed = False
|
||||
|
||||
async def __aenter__(self) -> "_FakeTTSStream":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb) -> bool: # noqa: ANN001
|
||||
return False
|
||||
|
||||
async def aclose(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
def raise_for_status(self) -> None:
|
||||
return None
|
||||
|
||||
async def aiter_bytes(self, chunk_size: int = 4096):
|
||||
self.chunk_size = chunk_size
|
||||
for chunk in self._chunks:
|
||||
yield chunk
|
||||
|
||||
|
||||
class _CaptureTTSClient:
|
||||
def __init__(self, chunks: list[bytes]) -> None:
|
||||
self.chunks = chunks
|
||||
self.calls: list[tuple[str, str, dict[str, object]]] = []
|
||||
|
||||
def stream(self, method: str, endpoint: str, *, json: dict[str, object]) -> _FakeTTSStream:
|
||||
self.calls.append((method, endpoint, dict(json)))
|
||||
return _FakeTTSStream(self.chunks)
|
||||
|
||||
|
||||
class VoiceServiceStreamTest(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_synthesize_stream_uses_payload_with_fake_client(self) -> None:
|
||||
client = _CaptureTTSClient([b"\x80\x80", b"\xff\x00"])
|
||||
service = VoiceService(api_key="test-key")
|
||||
service._client = client # type: ignore[assignment]
|
||||
voice = VoicePreset(
|
||||
preset="neutral",
|
||||
openai_voice="sage",
|
||||
rate=0.1,
|
||||
instructions="Keep the tone grounded.",
|
||||
)
|
||||
|
||||
chunks = [
|
||||
chunk
|
||||
async for chunk in service.synthesize_stream(
|
||||
"Spoken client reply",
|
||||
voice,
|
||||
response_format="opus",
|
||||
)
|
||||
]
|
||||
|
||||
self.assertEqual(len(client.calls), 1)
|
||||
method, endpoint, payload = client.calls[0]
|
||||
self.assertEqual(method, "POST")
|
||||
self.assertEqual(endpoint, TTS_ENDPOINT)
|
||||
self.assertEqual(payload["model"], TTS_MODEL)
|
||||
self.assertEqual(payload["voice"], "sage")
|
||||
self.assertEqual(payload["input"], "Spoken client reply")
|
||||
self.assertEqual(payload["response_format"], "opus")
|
||||
self.assertEqual(payload["speed"], 0.25)
|
||||
self.assertEqual(payload["instructions"], "Keep the tone grounded.")
|
||||
self.assertEqual([chunk.audio for chunk in chunks], [b"\x80\x80", b"\xff\x00"])
|
||||
|
||||
|
||||
class EndOfTurnDecisionTest(unittest.TestCase):
|
||||
def test_end_of_turn_requires_silence_threshold(self) -> None:
|
||||
pending = assess_end_of_turn(
|
||||
transcript_text="I am still talking",
|
||||
transcript_final=True,
|
||||
silence_ms=EOT_SILENCE_THRESHOLD_MS - 1,
|
||||
)
|
||||
ready = assess_end_of_turn(
|
||||
transcript_text="I am done",
|
||||
transcript_final=True,
|
||||
silence_ms=EOT_SILENCE_THRESHOLD_MS,
|
||||
)
|
||||
|
||||
self.assertFalse(pending.ready)
|
||||
self.assertTrue(pending.transcript_ready)
|
||||
self.assertFalse(pending.silence_ready)
|
||||
self.assertEqual(pending.reason, "silence_threshold_pending")
|
||||
self.assertTrue(ready.ready)
|
||||
self.assertEqual(ready.reason, "ready")
|
||||
|
||||
def test_end_of_turn_waits_for_final_transcript_even_after_silence(self) -> None:
|
||||
decision = assess_end_of_turn(
|
||||
transcript_text="Interim transcript",
|
||||
transcript_final=False,
|
||||
silence_ms=2500,
|
||||
)
|
||||
|
||||
self.assertFalse(decision.ready)
|
||||
self.assertFalse(decision.transcript_ready)
|
||||
self.assertTrue(decision.silence_ready)
|
||||
self.assertEqual(decision.reason, "final_transcript_pending")
|
||||
|
||||
def test_end_of_turn_rejects_empty_final_transcript(self) -> None:
|
||||
decision = assess_end_of_turn(
|
||||
transcript_text=" ",
|
||||
transcript_final=True,
|
||||
silence_ms=2500,
|
||||
)
|
||||
|
||||
self.assertFalse(decision.ready)
|
||||
self.assertFalse(decision.transcript_ready)
|
||||
self.assertTrue(decision.silence_ready)
|
||||
self.assertEqual(decision.reason, "empty_transcript")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
294
apps/api/app/test_voice_ws.py
Normal file
294
apps/api/app/test_voice_ws.py
Normal file
|
|
@ -0,0 +1,294 @@
|
|||
"""Deterministic WebSocket contract tests for the voice gateway route."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import unittest
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from .deps import Principal, Role
|
||||
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()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Loading…
Add table
Add a link
Reference in a new issue