세션 계약과 메모리 경계 보강

This commit is contained in:
Yun Chan 2026-06-28 23:52:18 +09:00
parent 391639c1de
commit 2bb052f624
12 changed files with 836 additions and 116 deletions

View file

@ -26,6 +26,7 @@ from .. import session_persistence
from ..deps import Principal, Role, require_role
from ..engine_client import EngineError, engine_client
from ..runtime_policy import runtime_fallback_allowed
from ..session_read_model import StageLabel, stage_label_or_none
from ..services import evaluator
from ..services.evaluator import SessionEvaluation, TurnEvaluation
from ..store import InProcSession
@ -50,11 +51,19 @@ class EvaluationSummary(BaseModel):
"""회기 평가 조회 응답(분포 + deep 결과 합본)."""
session_id: str
stage: str
stage: StageLabel | None = None
deep: Optional[dict[str, Any]] = None
distribution: dict[str, Any] = Field(default_factory=dict)
class TurnEvaluationResponse(TurnEvaluation):
stage: StageLabel
class SessionEvaluationResponse(SessionEvaluation):
stage: StageLabel
async def _load_session_or_404(session_id: str, principal: Principal) -> InProcSession:
sess = await session_persistence.load_session(session_id, principal, allow_ended=True)
if sess is not None:
@ -74,6 +83,10 @@ def _theory_mode_of(sess) -> Optional[str]:
return getattr(sess, "theory_mode", None)
def _summary_stage(value: object) -> StageLabel | None:
return stage_label_or_none(value)
@router.get("/health")
async def eval_health() -> dict[str, str]:
"""평가 라우터 헬스 — Features:evaluator 로 전환됨."""
@ -83,7 +96,7 @@ async def eval_health() -> dict[str, str]:
# ════════════════════════════════════════════════════════════════════════════
# 회기 deep-loop 재평가 트리거 (교수자/관리자)
# ════════════════════════════════════════════════════════════════════════════
@router.post("/sessions/{session_id}/reevaluate", response_model=SessionEvaluation)
@router.post("/sessions/{session_id}/reevaluate", response_model=SessionEvaluationResponse)
async def reevaluate_session(
session_id: str,
body: ReevaluateRequest,
@ -139,7 +152,7 @@ async def reevaluate_session(
# ════════════════════════════════════════════════════════════════════════════
# 단일 턴 fast-loop 재평가 트리거 (교수자/관리자)
# ════════════════════════════════════════════════════════════════════════════
@router.post("/sessions/{session_id}/turn", response_model=TurnEvaluation)
@router.post("/sessions/{session_id}/turn", response_model=TurnEvaluationResponse)
async def reevaluate_turn(
session_id: str,
body: TurnReevaluateRequest,
@ -211,13 +224,13 @@ async def get_session_evaluation(
await _load_session_or_404(session_id, principal)
record, _durable = await session_persistence.load_session_evaluation(session_id, principal)
if record is None:
return EvaluationSummary(session_id=session_id, stage="", deep=None, distribution={})
return EvaluationSummary(session_id=session_id, stage=None, deep=None, distribution={})
payload = record.get("payload")
deep = payload if isinstance(payload, dict) else {}
distribution = deep.get("distribution")
return EvaluationSummary(
session_id=session_id,
stage=str(record.get("stage") or deep.get("stage") or ""),
stage=_summary_stage(record.get("stage") or deep.get("stage")),
deep=deep,
distribution=distribution if isinstance(distribution, dict) else {},
)

View file

@ -121,7 +121,6 @@ _RECALL_CACHE: dict[str, memory.RecallContext] = {}
# 세션별 KB 증상 행동단서(회기 1회 산출·캐시). 빈 list 캐시 = 회기 내 재시도 안 함(안정성).
_KB_CUES_CACHE: dict[str, list[str]] = {}
_RAG_WARM_SEMAPHORE = asyncio.Semaphore(1)
_LEARNER_VISIBLE_AI_ROLE = "counselor"
# ────────────────────────────────────────────────────────────────────────────
# RAG 배선 헬퍼 — 내담자(CLIENT) 뷰. 임베더/KB/DB 풀 미가용 시 빈 값으로 graceful
@ -244,12 +243,6 @@ async def _ensure_kb_cues(session_id: str, card) -> list[str]:
return cues
async def _load_prev_case_summary(case_id: str) -> Optional[dict]:
"""직전 회기 요약(case 스코프) → build_recall_context 입력. 미존재/미가용 시 None."""
case_memory = await _load_case_memory(case_id)
return case_memory.get("prev_summary")
async def _load_case_memory(case_id: str) -> dict:
"""case-level 큰그림 + 직전 요약 + client-visible pinned fact를 한 번에 읽는다."""
empty = {"case_digest": None, "prev_summary": None, "pinned_facts": []}
@ -387,6 +380,30 @@ async def ensure_recall_context(sess: InProcSession) -> memory.RecallContext:
return recall
async def _prepare_turn_context(
*,
session_id: str,
sess: InProcSession,
learner_text: str,
) -> orchestrator.TurnContext:
recall = await ensure_recall_context(sess)
kb_cues = _KB_CUES_CACHE.get(session_id) or [] # 비차단: warm 전이면 빈 단서(graceful)
ctx = orchestrator.prepare_turn(
session_id=session_id,
case_id=sess.case_id,
card=sess.persona,
state=sess.state,
learner_text=learner_text,
recall_summary=recall.recall_summary,
pinned_facts=recall.pinned_facts,
recent_turns=sess.recent_turns(visible_to="client"),
kb_behavior_cues=kb_cues,
theory_mode=sess.theory_mode,
)
assert ctx.state_after is not None
return ctx
async def _warm_rag_caches(session_id: str, case_id: str, card) -> None:
"""RAG 회상·KB 행동단서를 **백그라운드**로 산출해 캐시한다(요청 경로 비차단).
@ -1032,22 +1049,11 @@ async def submit_turn(
"""Submit one trainee utterance and return the generated client reply."""
principal = _ensure_learner(principal)
sess = await _load_session_or_404(session_id, principal)
recall = await ensure_recall_context(sess)
kb_cues = _KB_CUES_CACHE.get(session_id) or [] # 비차단: warm 전이면 빈 단서(graceful)
ctx = orchestrator.prepare_turn(
ctx = await _prepare_turn_context(
session_id=session_id,
case_id=sess.case_id,
card=sess.persona,
state=sess.state,
learner_text=body.text,
recall_summary=recall.recall_summary,
pinned_facts=recall.pinned_facts,
recent_turns=sess.recent_turns(visible_to="client"),
kb_behavior_cues=kb_cues,
theory_mode=sess.theory_mode,
sess=sess,
)
assert ctx.state_after is not None
try:
result = await orchestrator.run_turn_generate(
@ -1159,22 +1165,11 @@ async def stream_turn(
"""Stream a generated client reply for one trainee utterance."""
principal = _ensure_learner(principal)
sess = await _load_session_or_404(session_id, principal)
recall = await ensure_recall_context(sess)
kb_cues = _KB_CUES_CACHE.get(session_id) or [] # 비차단: warm 전이면 빈 단서(graceful)
ctx = orchestrator.prepare_turn(
ctx = await _prepare_turn_context(
session_id=session_id,
case_id=sess.case_id,
card=sess.persona,
state=sess.state,
learner_text=body.text,
recall_summary=recall.recall_summary,
pinned_facts=recall.pinned_facts,
recent_turns=sess.recent_turns(visible_to="client"),
kb_behavior_cues=kb_cues,
theory_mode=sess.theory_mode,
sess=sess,
)
assert ctx.state_after is not None
async def event_generator():
last_beat = asyncio.get_running_loop().time()
@ -1229,7 +1224,7 @@ async def end_session(
session_id=session_id,
case_id=sess.case_id,
session_no=sess.session_no,
masked_turns=sess.masked_turns(),
masked_turns=sess.masked_turns(visible_to="client"),
prev_rapport_credit=sess.prev_rapport_credit,
open_threads=recall.open_threads,
)

View file

@ -13,9 +13,10 @@ from typing import Any
from fastapi import APIRouter, HTTPException, Request, Response, status
from fastapi.responses import HTMLResponse, PlainTextResponse
from pydantic import BaseModel, Field
from pydantic import BaseModel, Field, field_validator
from .. import session_persistence
from ..session_read_model import StageLabel, stage_label_or_none
router = APIRouter(tags=["share"])
@ -36,7 +37,7 @@ class PublicSessionShareResponse(BaseModel):
persona: str
date: str
durationLabel: str
reachedPhase: str
reachedPhase: StageLabel | None = None
sessionSignal: str
reviewReady: bool = False
goodMoments: list[str] = Field(default_factory=list)
@ -44,6 +45,11 @@ class PublicSessionShareResponse(BaseModel):
worksheetHighlights: list[dict[str, str]] = Field(default_factory=list)
privacy: str = ""
@field_validator("reachedPhase", mode="before")
@classmethod
def _normalize_reached_phase(cls, value: object) -> StageLabel | None:
return stage_label_or_none(value)
def _safe_payload(payload: dict[str, Any]) -> PublicSessionShareResponse:
return PublicSessionShareResponse(
@ -56,7 +62,7 @@ def _safe_payload(payload: dict[str, Any]) -> PublicSessionShareResponse:
persona=str(payload.get("persona") or ""),
date=str(payload.get("date") or ""),
durationLabel=str(payload.get("durationLabel") or ""),
reachedPhase=str(payload.get("reachedPhase") or ""),
reachedPhase=payload.get("reachedPhase"),
sessionSignal=str(payload.get("sessionSignal") or ""),
reviewReady=bool(payload.get("reviewReady")),
goodMoments=[str(item) for item in payload.get("goodMoments") or []][:3],

View file

@ -20,7 +20,7 @@ from __future__ import annotations
import re
from dataclasses import dataclass, field
from typing import Any, Callable, Optional
from typing import Any, Callable, Literal, Optional
from .state_machine import SessionState
@ -36,6 +36,7 @@ _COUNSELING_AGREEMENT_WITHDRAWAL_RE = re.compile(
r"(약속).*(취소|철회|못\s*지키|지키지\s*않|안\s*지키)|"
r"\s*이상.*(상담|회기).*(안\s*하|하지\s*않|못\s*하)"
)
_DIGEST_SPEAKERS = {"counselor", "client"}
# ════════════════════════════════════════════════════════════════════════════
@ -112,6 +113,42 @@ class CarryOver:
compression_job: Optional["CompressionJob"] = None
@dataclass(frozen=True, slots=True)
class MaskedDigestTurn:
"""A single client-visible, masked turn allowed into narrative digest input."""
speaker: str
text: str
turn_id: str | None = None
@dataclass(frozen=True, slots=True)
class SessionDigestInput:
"""Narrative-only digest boundary for future LLM compression.
This object intentionally excludes raw text, evaluation payloads, CCD, and
deterministic carry-over state. Numeric carry stays in CarryOver.end_state.
"""
session_id: str
case_id: str | None
session_no: int
masked_turns: tuple[MaskedDigestTurn, ...]
open_threads: tuple[str, ...] = ()
@dataclass(frozen=True, slots=True)
class SessionDigestResult:
"""Digest writer output contract shared by fallback and future LLM worker."""
session_id: str
case_id: str | None
session_no: int
digest: str
open_threads: tuple[str, ...] = ()
source: Literal["fallback", "llm"] = "fallback"
@dataclass(slots=True)
class CompressionJob:
"""회기종료 narrative 압축 작업(LLM, 비동기 비블로킹). 큐에 적재될 페이로드.
@ -120,12 +157,30 @@ class CompressionJob:
orchestrator/background task engine_client + RAG 수행한다(여기선 페이로드만).
"""
session_id: str
case_id: Optional[str]
session_no: int
masked_turns: list[dict[str, str]] # [{speaker, text}] (text_masked)
end_state: dict
open_threads: list[str] = field(default_factory=list)
digest_input: SessionDigestInput
@property
def session_id(self) -> str:
return self.digest_input.session_id
@property
def case_id(self) -> str | None:
return self.digest_input.case_id
@property
def session_no(self) -> int:
return self.digest_input.session_no
@property
def masked_turns(self) -> list[dict[str, str]]:
return [
{"speaker": turn.speaker, "text": turn.text}
for turn in self.digest_input.masked_turns
]
@property
def open_threads(self) -> list[str]:
return list(self.digest_input.open_threads)
@dataclass(frozen=True, slots=True)
@ -159,16 +214,71 @@ def make_carry_over(
end_state = state.snapshot()
rapport_delta = round(state.rapport_credit - prev_rapport_credit, 4)
job = CompressionJob(
digest_input=build_session_digest_input(
session_id=session_id,
case_id=case_id,
session_no=session_no,
masked_turns=masked_turns,
end_state=end_state,
open_threads=list(open_threads or []),
open_threads=open_threads,
),
)
return CarryOver(end_state=end_state, rapport_delta=rapport_delta, compression_job=job)
def _is_client_visible(turn: dict[str, Any]) -> bool:
visible_to = turn.get("visible_to")
if visible_to is None:
return True
if isinstance(visible_to, str):
return visible_to == "client"
try:
return "client" in {str(item) for item in visible_to}
except TypeError:
return False
def _masked_digest_turn(turn: dict[str, Any]) -> MaskedDigestTurn | None:
if not _is_client_visible(turn):
return None
speaker = str(turn.get("speaker") or "").strip()
if speaker not in _DIGEST_SPEAKERS:
return None
text = str(turn.get("text_masked") or turn.get("text") or "").strip()
if not text:
return None
turn_id = turn.get("turn_id")
return MaskedDigestTurn(
speaker=speaker,
text=text,
turn_id=str(turn_id) if turn_id else None,
)
def build_session_digest_input(
*,
session_id: str,
case_id: str | None,
session_no: int,
masked_turns: list[dict[str, Any]],
open_threads: Optional[list[str]] = None,
) -> SessionDigestInput:
"""Normalize the only payload shape allowed into narrative digest workers."""
turns = tuple(
digest_turn
for raw_turn in masked_turns
if (digest_turn := _masked_digest_turn(raw_turn)) is not None
)
threads = tuple(str(thread).strip() for thread in (open_threads or []) if str(thread).strip())
return SessionDigestInput(
session_id=session_id,
case_id=case_id,
session_no=int(session_no),
masked_turns=turns,
open_threads=threads,
)
def build_compression_messages(job: CompressionJob) -> list[dict[str, str]]:
"""CompressionJob → 서사 압축용 EngineMessage 평문(dict) 리스트.
@ -177,10 +287,10 @@ def build_compression_messages(job: CompressionJob) -> list[dict[str, str]]:
pinned 사실 보존·정답 미포함 지시 포함.
"""
transcript = "\n".join(
f"{('상담자' if t.get('speaker') == 'counselor' else '내담자')}: {t.get('text', '')}"
for t in job.masked_turns
f"{('상담자' if t.speaker == 'counselor' else '내담자')}: {t.text}"
for t in job.digest_input.masked_turns
)
threads = "\n".join(f"- {t}" for t in job.open_threads) or "(없음)"
threads = "\n".join(f"- {t}" for t in job.digest_input.open_threads) or "(없음)"
system = (
"당신은 상담 회기 종료 요약기다. 아래 마스킹된 축어록을 6~10문장 digest 로 압축한다.\n"
"규칙: ① 사실·정서 궤적·미해결 주제를 보존한다. ② 평가/점수/정답 라벨은 절대 포함하지 않는다.\n"
@ -188,7 +298,6 @@ def build_compression_messages(job: CompressionJob) -> list[dict[str, str]]:
)
user = (
f"[회기 번호] {job.session_no}\n"
f"[종료 상태(수치, 참고)] {job.end_state}\n"
f"[미해결 주제]\n{threads}\n\n"
f"[마스킹된 축어록]\n{transcript}\n\n"
"위를 digest 6~10문장으로 압축하라."
@ -214,8 +323,36 @@ def build_fallback_session_digest(
LLM 압축/embedding writer가 붙기 전에도 다음 회기 recall이 문자열로 남지 않도록
client-visible 마스킹 발화와 결정론 상태 수치만 사용한다.
"""
digest_input = build_session_digest_input(
session_id="",
case_id=None,
session_no=session_no,
masked_turns=masked_turns,
)
return build_fallback_digest_result(digest_input, end_state=end_state).digest
def build_fallback_digest_result(
digest_input: SessionDigestInput,
*,
end_state: dict,
) -> SessionDigestResult:
"""Build the current deterministic digest through the explicit contract."""
masked_turns = [
{"speaker": turn.speaker, "text": turn.text}
for turn in digest_input.masked_turns
]
session_no = digest_input.session_no
if not masked_turns:
return f"S{session_no}: 실제 발화가 없어 요약을 생성하지 않았다."
return SessionDigestResult(
session_id=digest_input.session_id,
case_id=digest_input.case_id,
session_no=session_no,
digest=f"S{session_no}: 실제 발화가 없어 요약을 생성하지 않았다.",
open_threads=digest_input.open_threads,
source="fallback",
)
counselor_count = sum(1 for turn in masked_turns if turn.get("speaker") == "counselor")
client_turns = [turn for turn in masked_turns if turn.get("speaker") == "client"]
@ -231,14 +368,23 @@ def build_fallback_session_digest(
status_bits.append(f"라포 {rapport}")
status = ", ".join(status_bits)
if last_client:
return (
digest = (
f"S{session_no}: 마스킹 축어록 기준 상담자 {counselor_count}회, "
f"내담자 {client_count}회 발화. 마지막 내담자 반응은 \"{last_client}\". {status}."
)
return (
else:
digest = (
f"S{session_no}: 마스킹 축어록 기준 상담자 {counselor_count}회, "
f"내담자 {client_count}회 발화. {status}."
)
return SessionDigestResult(
session_id=digest_input.session_id,
case_id=digest_input.case_id,
session_no=session_no,
digest=digest,
open_threads=digest_input.open_threads,
source="fallback",
)
def _fact_value(text: Any) -> str:
@ -369,8 +515,13 @@ __all__ = [
"build_recall_context",
"CarryOver",
"CompressionJob",
"MaskedDigestTurn",
"SessionDigestInput",
"SessionDigestResult",
"make_carry_over",
"build_session_digest_input",
"build_compression_messages",
"build_fallback_digest_result",
"build_fallback_session_digest",
"PinnedFactCandidate",
"extract_pinned_fact_candidates",

View file

@ -26,6 +26,12 @@ _SESSION_SHARE_TOKEN_INDEX: dict[str, str] = {}
_LIVE_COACH_EVENT_CACHE: dict[str, list[dict[str, Any]]] = {}
_SESSION_ARCHIVE_CACHE: dict[str, dict[str, Any]] = {}
_SESSION_AUDIT_ROLES = {"teacher", "admin"}
_WORKSHEET_REVIEW_STATUS_VALUES = {
"pending",
"approved",
"changes_requested",
"rejected",
}
_APPROPRIATENESS_SCORE = {
"warn": 1.0,
"neutral": 3.0,
@ -39,6 +45,17 @@ class CaseContext:
last_session_no: int
@dataclass(slots=True)
class SessionSummaryWrite:
session_id: str
case_id: str
session_no: int
end_state: dict[str, Any]
rapport_delta: float
digest: str
open_threads: list[str]
_JOINED_CARD_COLUMNS = (
"card_persona_id",
"card_code",
@ -935,11 +952,24 @@ async def ensure_review_tables() -> None:
status TEXT NOT NULL DEFAULT 'pending'
CHECK (status IN ('pending','viewed','closed')),
note TEXT NOT NULL DEFAULT '',
worksheet_status TEXT NOT NULL DEFAULT 'pending',
worksheet_note TEXT NOT NULL DEFAULT '',
worksheet_reviewed_at TIMESTAMPTZ,
reviewed_at TIMESTAMPTZ,
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
)
"""
)
await conn.execute(
"""
ALTER TABLE app.session_review_status
ADD COLUMN IF NOT EXISTS worksheet_status TEXT NOT NULL DEFAULT 'pending';
ALTER TABLE app.session_review_status
ADD COLUMN IF NOT EXISTS worksheet_note TEXT NOT NULL DEFAULT '';
ALTER TABLE app.session_review_status
ADD COLUMN IF NOT EXISTS worksheet_reviewed_at TIMESTAMPTZ;
"""
)
await conn.execute(
"""
ALTER TABLE app.session_review_status ENABLE ROW LEVEL SECURITY;
@ -1324,17 +1354,29 @@ def _review_status_from_row(row: Any) -> dict[str, Any]:
"reviewer_id": str(row["reviewer_id"] or ""),
"status": str(row["status"] or "pending"),
"note": str(row["note"] or ""),
"worksheet_status": _worksheet_review_status(row["worksheet_status"]),
"worksheet_note": str(row["worksheet_note"] or ""),
"worksheet_reviewed_at": _iso_dt(row["worksheet_reviewed_at"]),
"reviewed_at": _iso_dt(row["reviewed_at"]),
"updated_at": _iso_dt(row["updated_at"]),
}
def _worksheet_review_status(value: object) -> str:
raw = str(value or "pending")
if raw in _WORKSHEET_REVIEW_STATUS_VALUES:
return raw
return "pending"
def _review_status_cache_record(
*,
session_id: str,
reviewer_id: str,
status: str,
note: str,
worksheet_status: str | None = None,
worksheet_note: str | None = None,
) -> dict[str, Any]:
now = datetime.now(timezone.utc)
previous = _SESSION_REVIEW_STATUS_CACHE.get(session_id) or {}
@ -1343,11 +1385,25 @@ def _review_status_cache_record(
reviewed_at = now.isoformat().replace("+00:00", "Z")
if status != "closed":
reviewed_at = None
worksheet_reviewed_at = previous.get("worksheet_reviewed_at")
if worksheet_status is None:
worksheet_status = _worksheet_review_status(previous.get("worksheet_status"))
worksheet_note = str(previous.get("worksheet_note") or "")
else:
worksheet_status = _worksheet_review_status(worksheet_status)
worksheet_note = str(worksheet_note or "").strip()
if worksheet_status == "pending":
worksheet_reviewed_at = None
else:
worksheet_reviewed_at = now.isoformat().replace("+00:00", "Z")
return {
"session_id": session_id,
"reviewer_id": reviewer_id,
"status": status,
"note": note,
"worksheet_status": worksheet_status,
"worksheet_note": worksheet_note,
"worksheet_reviewed_at": worksheet_reviewed_at,
"reviewed_at": reviewed_at,
"updated_at": now.isoformat().replace("+00:00", "Z"),
}
@ -1368,7 +1424,8 @@ async def list_session_review_statuses(
) as conn:
rows = await conn.fetch(
"""
SELECT session_id, reviewer_id, status, note, reviewed_at, updated_at
SELECT session_id, reviewer_id, status, note, worksheet_status,
worksheet_note, worksheet_reviewed_at, reviewed_at, updated_at
FROM app.session_review_status
WHERE session_id = ANY($1::uuid[])
""",
@ -1402,14 +1459,20 @@ async def save_session_review_status(
status: str,
note: str,
principal: Principal,
worksheet_status: str | None = None,
worksheet_note: str | None = None,
) -> tuple[dict[str, Any] | None, bool]:
note = note.strip()
if worksheet_status is not None:
worksheet_note = str(worksheet_note or "").strip()
if runtime_fallback_allowed():
_SESSION_REVIEW_STATUS_CACHE[session_id] = _review_status_cache_record(
session_id=session_id,
reviewer_id=reviewer_id,
status=status,
note=note,
worksheet_status=worksheet_status,
worksheet_note=worksheet_note,
)
try:
get_pool()
@ -1421,10 +1484,17 @@ async def save_session_review_status(
row = await conn.fetchrow(
"""
INSERT INTO app.session_review_status (
session_id, reviewer_id, status, note, reviewed_at, updated_at
session_id, reviewer_id, status, note, worksheet_status,
worksheet_note, worksheet_reviewed_at, reviewed_at, updated_at
)
VALUES (
$1::uuid, $2::uuid, $3, $4,
COALESCE($5, 'pending'),
CASE WHEN $5::text IS NULL THEN '' ELSE COALESCE($6, '') END,
CASE
WHEN $5::text IS NULL OR $5 = 'pending' THEN NULL
ELSE now()
END,
CASE WHEN $3 = 'closed' THEN now() ELSE NULL END,
now()
)
@ -1432,18 +1502,38 @@ async def save_session_review_status(
reviewer_id = EXCLUDED.reviewer_id,
status = EXCLUDED.status,
note = EXCLUDED.note,
worksheet_status = CASE
WHEN $5::text IS NULL
THEN app.session_review_status.worksheet_status
ELSE $5
END,
worksheet_note = CASE
WHEN $5::text IS NULL
THEN app.session_review_status.worksheet_note
ELSE COALESCE($6, '')
END,
worksheet_reviewed_at = CASE
WHEN $5::text IS NULL
THEN app.session_review_status.worksheet_reviewed_at
WHEN $5 = 'pending'
THEN NULL
ELSE now()
END,
reviewed_at = CASE
WHEN EXCLUDED.status = 'closed'
THEN COALESCE(app.session_review_status.reviewed_at, now())
ELSE NULL
END,
updated_at = now()
RETURNING session_id, reviewer_id, status, note, reviewed_at, updated_at
RETURNING session_id, reviewer_id, status, note, worksheet_status,
worksheet_note, worksheet_reviewed_at, reviewed_at, updated_at
""",
session_id,
reviewer_id,
status,
note,
worksheet_status,
worksheet_note,
)
return (_review_status_from_row(row) if row else None), True
except Exception:
@ -2165,14 +2255,31 @@ async def _upsert_pinned_fact_candidates(conn: Any, sess: InProcSession) -> None
turn_id=fact.source_turn_id,
)
def _build_session_summary_write(sess: InProcSession, carry: memory.CarryOver) -> SessionSummaryWrite:
digest_input = memory.build_session_digest_input(
session_id=sess.session_id,
case_id=sess.case_id,
session_no=sess.session_no,
masked_turns=sess.masked_turns(visible_to="client"),
open_threads=carry.compression_job.open_threads if carry.compression_job else [],
)
digest_result = memory.build_fallback_digest_result(digest_input, end_state=carry.end_state)
return SessionSummaryWrite(
session_id=sess.session_id,
case_id=sess.case_id,
session_no=sess.session_no,
end_state=carry.end_state,
rapport_delta=carry.rapport_delta,
digest=digest_result.digest,
open_threads=list(digest_result.open_threads),
)
async def end_session(sess: InProcSession, carry: memory.CarryOver) -> bool:
try:
get_pool()
digest = memory.build_fallback_session_digest(
session_no=sess.session_no,
masked_turns=sess.masked_turns(visible_to="client"),
end_state=carry.end_state,
)
summary_write = _build_session_summary_write(sess, carry)
async with acquire(role="learner", user_id=sess.learner_id) as conn:
await conn.execute(
"""
@ -2196,13 +2303,13 @@ async def end_session(sess: InProcSession, carry: memory.CarryOver) -> bool:
digest = EXCLUDED.digest,
open_threads = EXCLUDED.open_threads
""",
sess.session_id,
sess.case_id,
sess.session_no,
carry.end_state,
carry.rapport_delta,
digest,
list(carry.compression_job.open_threads if carry.compression_job else []),
summary_write.session_id,
summary_write.case_id,
summary_write.session_no,
summary_write.end_state,
summary_write.rapport_delta,
summary_write.digest,
summary_write.open_threads,
)
case_row = await conn.fetchrow(
"""
@ -2222,7 +2329,7 @@ async def end_session(sess: InProcSession, carry: memory.CarryOver) -> bool:
case_digest = memory.merge_case_digest(
existing_digest=case_row["case_digest"],
session_no=sess.session_no,
session_digest=digest,
session_digest=summary_write.digest,
)
rapport_trajectory = memory.merge_rapport_trajectory(
case_row["rapport_trajectory"],

View file

@ -11,27 +11,26 @@ import re
from collections import Counter
from dataclasses import dataclass
from datetime import datetime
from typing import Literal, Optional
from typing import Literal, Optional, cast
from pydantic import BaseModel, Field
from . import turn_runtime
from .config import settings
from .services import session_metrics
from .stage_contract import (
ReviewPhaseKey,
StageLabel,
review_phase_key,
stage_label as _normalize_stage_label,
stage_label_or_none,
)
from .store import InProcSession, TurnRecord
StageLabel = Literal["라포", "탐색", "개입", "정리"]
WorksheetSpeaker = Literal["learner", "client"]
WorksheetItemSpec = tuple[str, str, list[str], WorksheetSpeaker | None]
WorksheetSectionSpec = tuple[str, str, list[WorksheetItemSpec]]
LEARNER_VISIBLE_AI_ROLE = "counselor"
_PHASE_KEY_BY_LABEL = {
"라포": "rapport",
"탐색": "explore",
"개입": "intervene",
"정리": "closing",
}
class LearnerSessionSummary(BaseModel):
@ -204,8 +203,8 @@ class ReviewTurn(BaseModel):
class ReviewPhaseSegment(BaseModel):
key: str
label: str
key: ReviewPhaseKey
label: StageLabel
weight: float
@ -330,6 +329,9 @@ class SessionTeacherReviewStatus(BaseModel):
reviewerId: str | None = None
reviewedAt: str | None = None
updatedAt: str | None = None
worksheetStatus: Literal["pending", "approved", "changes_requested", "rejected"] = "pending"
worksheetNote: str = ""
worksheetReviewedAt: str | None = None
class SessionReviewResponse(BaseModel):
@ -338,7 +340,7 @@ class SessionReviewResponse(BaseModel):
date: str
durationLabel: str
durationSeconds: int
reachedPhase: str
reachedPhase: StageLabel
sessionSignal: str
supervisorState: str
supervisorName: str
@ -385,8 +387,8 @@ class SessionReviewReadInput:
now_ts: float | None = None
def stage_label(stage: object) -> str:
return turn_runtime.stage_label(stage)
def stage_label(stage: object) -> StageLabel:
return cast(StageLabel, _normalize_stage_label(stage))
def iso(ts: float | None) -> str | None:
@ -658,7 +660,7 @@ def _client_name(raw: str) -> str:
return name or raw.strip() or "내담자"
def _review_summary(*, client_name: str, reached_phase: str, turns: list[ReviewTurn]) -> str:
def _review_summary(*, client_name: str, reached_phase: StageLabel, turns: list[ReviewTurn]) -> str:
if not turns:
return (
"아직 실제 발화가 없어 리뷰를 만들 수 없습니다. 회기를 진행한 뒤 종료하면 "
@ -674,11 +676,11 @@ def _review_summary(*, client_name: str, reached_phase: str, turns: list[ReviewT
)
def _phase_segments(stage_labels: list[str]) -> list[ReviewPhaseSegment]:
def _phase_segments(stage_labels: list[StageLabel]) -> list[ReviewPhaseSegment]:
counts = Counter(stage_labels)
return [
ReviewPhaseSegment(
key=_PHASE_KEY_BY_LABEL.get(label, label),
key=review_phase_key(label),
label=label,
weight=float(count),
)
@ -959,6 +961,15 @@ def saved_case_worksheet_from_payload(payload: dict[str, object] | None) -> Revi
return worksheet.model_copy(update={"status": "saved_by_learner"})
def _worksheet_review_status_value(
value: object,
) -> Literal["pending", "approved", "changes_requested", "rejected"]:
raw = str(value or "pending")
if raw in {"approved", "changes_requested", "rejected"}:
return raw # type: ignore[return-value]
return "pending"
def _worksheet_share_highlights(worksheet: ReviewCaseWorksheet, *, limit: int = 4) -> list[dict[str, str]]:
highlights: list[dict[str, str]] = []
for section in worksheet.sections:
@ -1342,6 +1353,9 @@ def build_session_review(read_input: SessionReviewReadInput) -> SessionReviewRes
reviewerId=str(review_status.get("reviewer_id") or "") or None,
reviewedAt=str(review_status.get("reviewed_at") or "") or None,
updatedAt=str(review_status.get("updated_at") or "") or None,
worksheetStatus=_worksheet_review_status_value(review_status.get("worksheet_status")),
worksheetNote=str(review_status.get("worksheet_note") or ""),
worksheetReviewedAt=str(review_status.get("worksheet_reviewed_at") or "") or None,
)
return SessionReviewResponse(

View file

@ -0,0 +1,54 @@
"""Shared stage labels for browser-facing session contracts.
State-machine internals may use enum values, legacy code values, or persisted
strings. API DTOs should expose only these Korean labels, and compatibility
normalization belongs in one import-light module.
"""
from __future__ import annotations
from typing import Literal, cast
StageLabel = Literal["라포", "탐색", "개입", "정리"]
ReviewPhaseKey = Literal["rapport", "explore", "intervene", "closing"]
STAGE_LABEL_VALUES: tuple[StageLabel, ...] = ("라포", "탐색", "개입", "정리")
_STAGE_LABEL_BY_CODE: dict[str, StageLabel] = {
"RAPPORT": "라포",
"EXPLORE": "탐색",
"INTERVENE": "개입",
"CLOSE": "정리",
"rapport": "라포",
"explore": "탐색",
"intervene": "개입",
"close": "정리",
}
_PHASE_KEY_BY_LABEL: dict[StageLabel, ReviewPhaseKey] = {
"라포": "rapport",
"탐색": "explore",
"개입": "intervene",
"정리": "closing",
}
def stage_label(stage: object) -> StageLabel | str:
"""Normalize Stage enum/string inputs to the browser-facing Korean label."""
name = getattr(stage, "name", "")
raw = str(getattr(stage, "value", stage))
return _STAGE_LABEL_BY_CODE.get(name) or _STAGE_LABEL_BY_CODE.get(raw) or raw
def stage_label_or_none(stage: object) -> StageLabel | None:
"""Normalize known stages and drop blank or unknown legacy values."""
label = stage_label(stage)
if label in STAGE_LABEL_VALUES:
return cast(StageLabel, label)
return None
def review_phase_key(label: StageLabel) -> ReviewPhaseKey:
return _PHASE_KEY_BY_LABEL[label]

View file

@ -45,6 +45,90 @@ class SessionMemoryPureTest(unittest.TestCase):
self.assertIn("[ORG]", digest)
self.assertNotIn("김서연", digest)
def test_session_digest_input_keeps_only_client_visible_masked_turns(self) -> None:
digest_input = memory.build_session_digest_input(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
session_no=4,
masked_turns=[
{
"speaker": "client",
"text": "raw 김서연",
"text_masked": "저는 [NAME]입니다.",
"turn_id": "00000000-0000-0000-0000-000000000201",
"visible_to": ["client", "evaluator"],
},
{
"speaker": "client",
"text": "평가자 전용 raw",
"text_masked": "평가자 전용 masked",
"visible_to": ["evaluator"],
},
{
"speaker": "system",
"text": "시스템 메모",
"visible_to": ["client"],
},
],
open_threads=[" 가족 이야기 이어가기 ", ""],
)
self.assertEqual(digest_input.session_no, 4)
self.assertEqual(len(digest_input.masked_turns), 1)
self.assertEqual(digest_input.masked_turns[0].speaker, "client")
self.assertEqual(digest_input.masked_turns[0].text, "저는 [NAME]입니다.")
self.assertEqual(
digest_input.masked_turns[0].turn_id,
"00000000-0000-0000-0000-000000000201",
)
self.assertEqual(digest_input.open_threads, ("가족 이야기 이어가기",))
self.assertNotIn("김서연", " ".join(turn.text for turn in digest_input.masked_turns))
def test_compression_messages_use_digest_contract_without_end_state(self) -> None:
digest_input = memory.build_session_digest_input(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
session_no=5,
masked_turns=[
{"speaker": "client", "text": "저는 [NAME]입니다.", "visible_to": ["client"]},
{"speaker": "counselor", "text": "그때 마음을 더 말해볼까요?", "visible_to": ["client"]},
],
open_threads=["다음 회기에 가족 이야기를 이어가기"],
)
messages = memory.build_compression_messages(memory.CompressionJob(digest_input=digest_input))
prompt = "\n".join(message["content"] for message in messages)
self.assertIn("저는 [NAME]입니다.", prompt)
self.assertIn("다음 회기에 가족 이야기를 이어가기", prompt)
self.assertNotIn("김서연", prompt)
self.assertNotIn("종료 상태", prompt)
self.assertNotIn("rapport_credit", prompt)
self.assertNotIn("evaluation", prompt)
def test_fallback_digest_result_uses_shared_digest_contract(self) -> None:
digest_input = memory.build_session_digest_input(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
session_no=6,
masked_turns=[
{"speaker": "client", "text": "저는 [NAME]입니다.", "visible_to": ["client"]},
],
open_threads=["가족 이야기"],
)
result = memory.build_fallback_digest_result(
digest_input,
end_state={"stage": "탐색", "effective_openness": 0.4, "rapport_credit": 0.3},
)
self.assertEqual(result.source, "fallback")
self.assertEqual(result.session_id, "00000000-0000-0000-0000-00000000feed")
self.assertEqual(result.case_id, "00000000-0000-0000-0000-00000000ca5e")
self.assertEqual(result.session_no, 6)
self.assertEqual(result.open_threads, ("가족 이야기",))
self.assertIn("S6:", result.digest)
self.assertIn("[NAME]", result.digest)
self.assertNotIn("김서연", result.digest)
def test_extract_pinned_fact_candidates_is_conservative_and_masked(self) -> None:
facts = memory.extract_pinned_fact_candidates(
[
@ -363,13 +447,35 @@ class SessionMemoryPersistenceTest(unittest.IsolatedAsyncioTestCase):
persisted = await session_persistence.end_session(sess, carry)
self.assertTrue(persisted)
summary_writes = [
args for query, args in conn.executed if "INSERT INTO app.session_summary" in query
]
self.assertEqual(len(summary_writes), 1)
(
summary_session_id,
summary_case_id,
summary_session_no,
summary_end_state,
summary_rapport_delta,
summary_digest,
summary_open_threads,
) = summary_writes[0]
self.assertEqual(summary_session_id, sess.session_id)
self.assertEqual(summary_case_id, case_id)
self.assertEqual(summary_session_no, 2)
self.assertEqual(summary_end_state, carry.end_state)
self.assertEqual(summary_rapport_delta, 0.21)
self.assertTrue(summary_digest.startswith("S2:"))
self.assertIn("[NAME]", summary_digest)
self.assertNotIn("김서연", summary_digest)
self.assertEqual(summary_open_threads, [])
case_updates = [
args for query, args in conn.executed if "UPDATE app.case_profile" in query
]
self.assertEqual(len(case_updates), 1)
_, _, case_digest, trajectory, alliance_level = case_updates[0]
self.assertIn("S1: 이전 회기", case_digest)
self.assertIn("S2:", case_digest)
self.assertIn(summary_digest, case_digest)
self.assertIn("[NAME]", case_digest)
self.assertNotIn("김서연", case_digest)
self.assertEqual(trajectory[-1]["session_no"], 2)

View file

@ -103,9 +103,43 @@ class SessionShareTest(unittest.IsolatedAsyncioTestCase):
summary = await share_routes._load_share_or_404(token)
self.assertIn("Vignette 회기 리뷰", summary.title)
self.assertIn(summary.reachedPhase, {"라포", "탐색", "개입", "정리"})
self.assertNotIn("원문 학습자 민감 발화", summary.model_dump_json())
self.assertIn("저장된 실제 축어록 2개", summary.summary)
async def test_public_share_normalizes_legacy_reached_phase(self) -> None:
coded = share_routes._safe_payload({"reachedPhase": "rapport"})
blank = share_routes._safe_payload({"reachedPhase": ""})
invalid = share_routes._safe_payload({"reachedPhase": "unknown-stage"})
missing = share_routes._safe_payload({})
self.assertEqual(coded.reachedPhase, "라포")
self.assertIsNone(blank.reachedPhase)
self.assertIsNone(invalid.reachedPhase)
self.assertIsNone(missing.reachedPhase)
async def test_public_share_lookup_tolerates_legacy_blank_reached_phase(self) -> None:
token = "legacyShareTokenValue000000000000"
token_hash = session_persistence.share_token_hash(token)
session_persistence._SESSION_SHARE_CACHE["legacy-share-session"] = {
"session_id": "legacy-share-session",
"token_hash": token_hash,
"payload": {
"title": "Legacy share",
"description": "Legacy summary",
"reachedPhase": "",
},
"created_at": 1_800_000_000.0,
"updated_at": 1_800_000_000.0,
"revoked_at": None,
}
session_persistence._SESSION_SHARE_TOKEN_INDEX[token_hash] = "legacy-share-session"
summary = await share_routes._load_share_or_404(token)
self.assertEqual(summary.title, "Legacy share")
self.assertIsNone(summary.reachedPhase)
async def test_revoke_share_blocks_public_lookup(self) -> None:
principal = _principal()
_ended_session(principal)

View file

@ -72,9 +72,13 @@ async def _consume_event_source(response: object) -> bytes:
class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self) -> None:
store._sessions.clear()
sessions._RECALL_CACHE.clear()
sessions._KB_CUES_CACHE.clear()
async def asyncTearDown(self) -> None:
store._sessions.clear()
sessions._RECALL_CACHE.clear()
sessions._KB_CUES_CACHE.clear()
async def test_append_turn_writes_provider_events_to_db(self) -> None:
class FakeConn:
@ -157,9 +161,16 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
async def test_generate_turn_persists_client_engine_telemetry(self) -> None:
principal = _principal()
sess = _session(principal)
sessions._RECALL_CACHE[sess.session_id] = memory.RecallContext(
recall_summary="직전 회기에서 김서연은 가족 이야기를 열어두었다.",
pinned_facts=["박민수와 주 1회 상담 약속"],
)
async def successful_turn(ctx, engine, **kwargs):
assert ctx.state_after is not None
self.assertIn("[NAME]", ctx.recall_summary or "")
self.assertNotIn("김서연", ctx.recall_summary or "")
self.assertEqual(ctx.pinned_facts, ["[NAME]와 주 1회 상담 약속"])
return orchestrator.TurnResult(
turn_seq=ctx.state_after.turn_seq,
stage=ctx.state_after.stage.value,
@ -429,8 +440,41 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
self.assertEqual(history.source, "runtime")
self.assertEqual(len(history.events), 1)
self.assertEqual(history.events[0].turn_seq, 1)
self.assertEqual(history.events[0].stage, "라포")
self.assertEqual(history.events[0].suggestion.title, response.title)
self.assertIn("학교", history.events[0].learner_text_excerpt or "")
session_persistence._LIVE_COACH_EVENT_CACHE[sess.session_id][0]["stage"] = "unknown-stage"
legacy_history = await sessions.list_live_coach_history(sess.session_id, principal)
self.assertIsNone(legacy_history.events[0].stage)
async def test_live_coach_event_normalizes_legacy_stage_values(self) -> None:
suggestion = live_coach.LiveCoachSuggestion(
status="degraded",
tone="neutral",
focus="exploration",
title="코칭",
message="다음 발화를 준비하세요.",
)
coded = live_coach.LiveCoachEvent(
event_id="event-1",
session_id="session-1",
turn_seq=1,
stage="rapport",
created_at="2026-06-28T00:00:00Z",
suggestion=suggestion,
)
invalid = live_coach.LiveCoachEvent(
event_id="event-2",
session_id="session-1",
turn_seq=2,
stage="unknown-stage",
created_at="2026-06-28T00:00:00Z",
suggestion=suggestion,
)
self.assertEqual(coded.stage, "라포")
self.assertIsNone(invalid.stage)
async def test_live_coach_uses_official_risk_reference_pack_for_crisis_signal(self) -> None:
item = live_coach.LiveCoachInput(
@ -523,6 +567,124 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
self.assertEqual(response.recall_summary, recall.recall_summary)
self.assertIs(sessions._RECALL_CACHE[response.session_id], recall)
async def test_next_session_turn_injects_seed_recall_into_engine_messages(self) -> None:
principal = _principal()
card = persona_service.P1
case_context = sessions.session_persistence.CaseContext(
case_id="00000000-0000-0000-0000-00000000ca5e",
last_session_no=1,
)
catalog_persona = SimpleNamespace(
card=card,
persona_id="00000000-0000-0000-0000-0000000000a1",
version=3,
degraded=False,
)
test_case = self
class FakeConn:
async def fetchrow(self, query: str, *args: object):
if "FROM app.case_profile" in query:
return {"case_digest": "S1: 김서연은 가족 이야기를 열어두었다."}
if "FROM app.session_summary" in query:
return {
"digest": "직전 회기에서 김서연은 침묵 이후 학교 이야기를 꺼냈다.",
"open_threads": ["다음 회기에서 상담 지속 의사를 확인하기"],
"end_state": {
"rapport_credit": 0.55,
"resistance": card.base_resistance(),
"ideation_stage": card.ideation_baseline(),
},
}
return None
async def fetch(self, query: str, *args: object):
test_case.assertIn("FROM app.pinned_fact", query)
test_case.assertIn("$2 = ANY(visible_to)", query)
return [{"value": "박민수와 주 1회 상담 약속"}]
class FakeAcquire:
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
async def fake_create_session(**kwargs):
return InProcSession(
session_id="db-seed-recall-session",
case_id=kwargs["case_id"],
learner_id=principal.user_id,
persona_code=card.code,
theory_mode=kwargs["theory_mode"],
persona=card,
state=kwargs["state"],
session_no=kwargs["session_no"],
prev_rapport_credit=kwargs["carry_rapport"],
)
def close_background(coro):
coro.close()
return None
with patch.object(sessions, "get_catalog_persona", AsyncMock(return_value=catalog_persona)), patch.object(
sessions.session_persistence,
"get_case_context",
AsyncMock(return_value=case_context),
), patch.object(sessions.db, "get_pool", return_value=object()), patch.object(
sessions.db,
"acquire",
return_value=FakeAcquire(FakeConn()),
), patch.object(
sessions.session_persistence,
"create_session",
fake_create_session,
), patch.object(sessions.asyncio, "create_task", close_background):
response = await sessions.start_session(
sessions.SessionStartRequest(persona_code=card.code),
principal,
)
started = store.get(response.session_id)
self.assertIsNotNone(started)
captured_messages: list[str] = []
async def successful_turn(ctx, engine, **kwargs):
assert ctx.state_after is not None
captured_messages.extend(message.content for message in ctx.messages)
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,
)
with patch.object(sessions, "_load_session_or_404", AsyncMock(return_value=started)), patch.object(
sessions.orchestrator,
"run_turn_generate",
successful_turn,
):
await sessions.submit_turn(
response.session_id,
sessions.TurnRequest(text="지난번 이야기를 이어가도 괜찮을까요?"),
principal,
)
message_blob = "\n".join(captured_messages)
self.assertIn("[L2 회상", message_blob)
self.assertIn("[케이스 큰그림]", message_blob)
self.assertIn("[직전 회기 요약]", message_blob)
self.assertIn("다음 회기에서 상담 지속 의사를 확인하기", message_blob)
self.assertIn("[L4 고정 사실", message_blob)
self.assertIn("[NAME]와 주 1회 상담 약속", message_blob)
self.assertNotIn("김서연", message_blob)
self.assertNotIn("박민수", message_blob)
async def test_start_session_requires_learner_consent_before_catalog_lookup(self) -> None:
principal = _principal()
principal.consent_at = None
@ -960,6 +1122,75 @@ class SessionTurnPersistenceTest(unittest.IsolatedAsyncioTestCase):
self.assertEqual(worksheet.limitations, ["학습자 저장본"])
self.assertEqual(worksheet.savedAt, "2026-06-27T10:00:00+00:00")
async def test_teacher_review_includes_manual_worksheet_decision(self) -> None:
learner = _principal()
teacher_principal = Principal(
user_id="00000000-0000-0000-0000-000000000902",
role=Role.TEACHER,
)
sess = _session(learner)
sess.ended = True
sess.ended_at = sess.created_at + 600
sess.turns.append(
TurnRecord(
turn_seq=1,
speaker="client",
stage=sess.state.stage.value,
text="저장본 검수를 확인합니다.",
text_masked="저장본 검수를 확인합니다.",
created_at=sess.created_at + 1,
)
)
saved_payload = {
"status": "saved_by_learner",
"generatedBy": "learner-edited worksheet",
"sections": [],
"limitations": ["학습자 저장본"],
"savedAt": "2026-06-27T10:00:00+00:00",
}
review_status = {
"session_id": sess.session_id,
"reviewer_id": teacher_principal.user_id,
"status": "viewed",
"note": "회기 전체 검토 메모",
"worksheet_status": "changes_requested",
"worksheet_note": "주호소 근거를 더 명확히 쓰도록 지도",
"worksheet_reviewed_at": "2026-06-27T10:05:00Z",
"reviewed_at": "",
"updated_at": "2026-06-27T10:05:00Z",
}
with (
patch.object(
sessions.session_persistence,
"load_session",
AsyncMock(return_value=sess),
),
patch.object(
sessions.session_persistence,
"load_case_worksheet",
AsyncMock(return_value=(saved_payload, True)),
),
patch.object(
sessions.session_persistence,
"load_session_evaluation",
AsyncMock(return_value=(None, False)),
),
patch.object(
sessions.session_persistence,
"load_session_review_status",
AsyncMock(return_value=(review_status, True)),
),
):
response = await sessions.get_session_review(sess.session_id, teacher_principal)
self.assertIsNotNone(response.teacherReview)
assert response.teacherReview is not None
self.assertEqual(response.caseWorksheet.status, "saved_by_learner")
self.assertEqual(response.teacherReview.worksheetStatus, "changes_requested")
self.assertIn("주호소", response.teacherReview.worksheetNote)
self.assertEqual(response.teacherReview.worksheetReviewedAt, "2026-06-27T10:05:00Z")
async def test_learner_can_save_case_formulation_worksheet(self) -> None:
principal = _principal()
sess = _session(principal)

View file

@ -0,0 +1,26 @@
from __future__ import annotations
import unittest
from . import stage_contract
from .services.state_machine import Stage
class StageContractTests(unittest.TestCase):
def test_stage_label_normalizes_enum_and_legacy_code_values(self) -> None:
self.assertEqual(stage_contract.stage_label(Stage.RAPPORT), "라포")
self.assertEqual(stage_contract.stage_label("rapport"), "라포")
self.assertEqual(stage_contract.stage_label("close"), "정리")
def test_stage_label_or_none_drops_unknown_legacy_values(self) -> None:
self.assertEqual(stage_contract.stage_label_or_none("intervene"), "개입")
self.assertIsNone(stage_contract.stage_label_or_none(""))
self.assertIsNone(stage_contract.stage_label_or_none("unknown-stage"))
def test_review_phase_key_uses_browser_contract_keys(self) -> None:
self.assertEqual(stage_contract.review_phase_key("라포"), "rapport")
self.assertEqual(stage_contract.review_phase_key("정리"), "closing")
if __name__ == "__main__":
unittest.main()

View file

@ -14,26 +14,9 @@ from . import db, session_persistence
from .deps import Principal
from .runtime_policy import require_runtime_fallback_allowed, runtime_fallback_allowed
from .services import guardrail, orchestrator, state_machine
from .stage_contract import stage_label
from .store import InProcSession, TurnRecord, store
_STAGE_LABELS = {
"RAPPORT": "라포",
"EXPLORE": "탐색",
"INTERVENE": "개입",
"CLOSE": "정리",
"rapport": "라포",
"explore": "탐색",
"intervene": "개입",
"close": "정리",
}
def stage_label(stage: object) -> str:
"""Stage enum과 문자열 값을 같은 한글 라벨로 정규화한다."""
name = getattr(stage, "name", "")
raw = str(getattr(stage, "value", stage))
return _STAGE_LABELS.get(name) or _STAGE_LABELS.get(raw) or raw
class SessionAccessError(str, Enum):
NOT_FOUND = "not_found"