532 lines
20 KiB
Python
532 lines
20 KiB
Python
"""회기 라이프사이클 메모리 — 시작 회상 + 종료 carry-over (4계층 매핑).
|
||
|
||
MEMORY_KNOWLEDGE_PERSONA_DESIGN §1·§2·§8 + 4대 대원칙(P1~P4):
|
||
① WORKING : 상태머신 수치(state_machine.SessionState) — 매 턴 체크포인트
|
||
② EPISODIC : 발화(turns) append-only — 회상 검색 대상
|
||
③ SUMMARY : 회기종료 압축(end_state 무손실 + digest 서사 LLM 압축)
|
||
④ SEMANTIC : case_profile 누적 + pinned_fact
|
||
|
||
핵심 원칙:
|
||
- P2/P4: 숫자(상태 수치)는 코드가 무손실 복사(carry-over). 서사(narrative)만 LLM 압축.
|
||
- P3: 회상은 큰그림→세부 순서(case_digest → 직전 summary → episodic recall) 토큰 예산 배분.
|
||
- 회상 요약(recall_summary)에는 CCD·정답·평가가 절대 들어가지 않는다(M6, 내담자 뷰).
|
||
|
||
이 모듈은 *순수 조립/룰 로직* + (선택) LLM 압축 *트리거 큐*만 담당한다.
|
||
실제 임베딩/하이브리드 검색은 RAG(Features) 소유 → 여기선 인터페이스(주입형)로 추상화한다.
|
||
DB 미가용(Docker off) 시에도 동작하도록 입력은 plain dict/list 로 받는다.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import re
|
||
from dataclasses import dataclass, field
|
||
from typing import Any, Callable, Literal, Optional
|
||
|
||
from .state_machine import SessionState
|
||
|
||
|
||
_SESSION_DIGEST_EXCERPT_CHARS = 90
|
||
_CASE_DIGEST_MAX_ENTRIES = 12
|
||
_RAPPORT_TRAJECTORY_MAX_ENTRIES = 24
|
||
_PINNED_FACT_MAX_VALUE_CHARS = 160
|
||
_COUNSELING_AGREEMENT_RE = re.compile(r"(상담|회기).*(주\s*\d+\s*회|매주|약속|계속|이어)|"
|
||
r"(주\s*\d+\s*회|매주).*(상담|회기)")
|
||
_COUNSELING_AGREEMENT_WITHDRAWAL_RE = re.compile(
|
||
r"(상담|회기).*(그만|중단|취소|철회|안\s*하|하지\s*않|못\s*하|이어\s*가지\s*않|계속\s*하지\s*않)|"
|
||
r"(약속).*(취소|철회|못\s*지키|지키지\s*않|안\s*지키)|"
|
||
r"더\s*이상.*(상담|회기).*(안\s*하|하지\s*않|못\s*하)"
|
||
)
|
||
_DIGEST_SPEAKERS = {"counselor", "client"}
|
||
|
||
|
||
# ════════════════════════════════════════════════════════════════════════════
|
||
# 회기 시작 — 회상 (큰그림 → 세부)
|
||
# ════════════════════════════════════════════════════════════════════════════
|
||
@dataclass(slots=True)
|
||
class RecallContext:
|
||
"""회기 시작 회상 결과. 내담자 AI system L2 주입용.
|
||
|
||
⚠️ CCD/정답/평가 미포함(내담자 뷰). digest/open_threads 는 "자기 기억" 표면만.
|
||
"""
|
||
|
||
recall_summary: Optional[str] = None # UI 카드 + L2 주입(큰그림→세부 합본)
|
||
pinned_facts: list[str] = field(default_factory=list) # L4 hard-pin(무손실)
|
||
open_threads: list[str] = field(default_factory=list)
|
||
carry: Optional[dict] = None # 직전 end_state(결정론 carry-over 입력)
|
||
|
||
|
||
def build_recall_context(
|
||
*,
|
||
case_digest: Optional[str] = None, # ④ 큰그림 (~400토큰 예산)
|
||
prev_summary: Optional[dict] = None, # ③ 직전 session_summary {digest, open_threads, homework, end_state}
|
||
episodic_snippets: Optional[list[str]] = None, # ② recall top-k 세부 (RAG 주입형)
|
||
pinned_facts: Optional[list[str]] = None, # ④ pinned_fact value[]
|
||
) -> RecallContext:
|
||
"""회상 컨텍스트 조립 (P3: 큰그림→세부 순서로 토큰 예산 배분).
|
||
|
||
DB/RAG 가 없으면 인자들이 None → 빈 RecallContext(첫 회기·in-proc fallback).
|
||
호출부(sessions.start)가 DB/RAG 가용 시 채워 넣는다.
|
||
"""
|
||
lines: list[str] = []
|
||
if case_digest:
|
||
lines.append(f"[케이스 큰그림]\n{case_digest}")
|
||
prev_end_state: Optional[dict] = None
|
||
open_threads: list[str] = []
|
||
if prev_summary:
|
||
digest = prev_summary.get("digest")
|
||
if digest:
|
||
lines.append(f"[직전 회기 요약]\n{digest}")
|
||
open_threads = list(prev_summary.get("open_threads") or [])
|
||
if open_threads:
|
||
ot = "\n".join(f"- {t}" for t in open_threads)
|
||
lines.append(f"[미해결 주제]\n{ot}")
|
||
homework = prev_summary.get("homework")
|
||
if homework:
|
||
lines.append(f"[지난 과제]\n{homework}")
|
||
prev_end_state = prev_summary.get("end_state")
|
||
if episodic_snippets:
|
||
snips = "\n".join(f"- {s}" for s in episodic_snippets)
|
||
lines.append(f"[지난 대화 단편(세부)]\n{snips}")
|
||
|
||
recall_summary = "\n\n".join(lines) if lines else None
|
||
return RecallContext(
|
||
recall_summary=recall_summary,
|
||
pinned_facts=list(pinned_facts or []),
|
||
open_threads=open_threads,
|
||
carry=prev_end_state,
|
||
)
|
||
|
||
|
||
# ════════════════════════════════════════════════════════════════════════════
|
||
# 회기 종료 — carry-over (무손실 수치 복사 + 서사 압축 트리거)
|
||
# ════════════════════════════════════════════════════════════════════════════
|
||
@dataclass(slots=True)
|
||
class CarryOver:
|
||
"""회기 종료 무손실 carry-over (P4: 코드 복사, LLM 미경유).
|
||
|
||
end_state = 상태머신 종료 snapshot(다음 회기 init_state 입력).
|
||
compression_job = 서사 digest LLM 압축이 *필요한* 입력 묶음(비동기 큐 대상).
|
||
"""
|
||
|
||
end_state: dict
|
||
rapport_delta: float = 0.0
|
||
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, 비동기 비블로킹). 큐에 적재될 페이로드.
|
||
|
||
입력은 *마스킹된 발화*만(F-03). 실제 LLM 호출/임베딩/DB UPSERT 는
|
||
orchestrator/background task 가 engine_client + RAG 로 수행한다(여기선 페이로드만).
|
||
"""
|
||
|
||
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)
|
||
class PinnedFactCandidate:
|
||
"""Rule-derived client-visible fact candidate for app.pinned_fact."""
|
||
|
||
key: str
|
||
value: str
|
||
fact_type: str
|
||
status: str = "stable"
|
||
confidence: float = 0.7
|
||
source_turn_id: str | None = None
|
||
|
||
|
||
def make_carry_over(
|
||
*,
|
||
state: SessionState,
|
||
session_id: str,
|
||
case_id: Optional[str],
|
||
session_no: int,
|
||
masked_turns: list[dict[str, str]],
|
||
prev_rapport_credit: float = 0.0,
|
||
open_threads: Optional[list[str]] = None,
|
||
) -> CarryOver:
|
||
"""회기 종료 carry-over 생성.
|
||
|
||
(A) 무손실: end_state = state.snapshot() (코드 복사) [P4]
|
||
(B) rapport_delta = 종료 rapport_credit − 이전 회기 rapport_credit
|
||
(C) 서사 압축은 CompressionJob 으로 큐잉(LLM, 비동기) — 여기선 페이로드만 만든다
|
||
"""
|
||
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,
|
||
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) 리스트.
|
||
|
||
실제 호출은 orchestrator/background 가 engine_client.generate(GenerateRequest(
|
||
ai_role='evaluator', tier='feedback', ...)) 로 수행. 여기선 프롬프트만 조립(IO 없음).
|
||
pinned 사실 보존·정답 미포함 지시 포함.
|
||
"""
|
||
transcript = "\n".join(
|
||
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.digest_input.open_threads) or "(없음)"
|
||
system = (
|
||
"당신은 상담 회기 종료 요약기다. 아래 마스킹된 축어록을 6~10문장 digest 로 압축한다.\n"
|
||
"규칙: ① 사실·정서 궤적·미해결 주제를 보존한다. ② 평가/점수/정답 라벨은 절대 포함하지 않는다.\n"
|
||
"③ 내담자가 실제로 말한 사실은 바꾸지 않는다(무손실). ④ 한국어, 간결한 임상 서술체."
|
||
)
|
||
user = (
|
||
f"[회기 번호] {job.session_no}\n"
|
||
f"[미해결 주제]\n{threads}\n\n"
|
||
f"[마스킹된 축어록]\n{transcript}\n\n"
|
||
"위를 digest 6~10문장으로 압축하라."
|
||
)
|
||
return [{"role": "system", "content": system}, {"role": "user", "content": user}]
|
||
|
||
|
||
def _compact(value: Any, *, limit: int = _SESSION_DIGEST_EXCERPT_CHARS) -> str:
|
||
text = " ".join(str(value or "").split())
|
||
if len(text) <= limit:
|
||
return text
|
||
return text[: max(0, limit - 3)].rstrip() + "..."
|
||
|
||
|
||
def build_fallback_session_digest(
|
||
*,
|
||
session_no: int,
|
||
masked_turns: list[dict[str, str]],
|
||
end_state: dict,
|
||
) -> str:
|
||
"""마스킹 축어록 기반 임시 회기 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 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"]
|
||
client_count = len(client_turns)
|
||
last_client = _compact(client_turns[-1].get("text") if client_turns else "")
|
||
stage = str(end_state.get("stage") or "미확인")
|
||
openness = end_state.get("effective_openness")
|
||
rapport = end_state.get("rapport_credit")
|
||
status_bits = [f"종료 단계 {stage}"]
|
||
if openness is not None:
|
||
status_bits.append(f"개방도 {openness}")
|
||
if rapport is not None:
|
||
status_bits.append(f"라포 {rapport}")
|
||
status = ", ".join(status_bits)
|
||
if last_client:
|
||
digest = (
|
||
f"S{session_no}: 마스킹 축어록 기준 상담자 {counselor_count}회, "
|
||
f"내담자 {client_count}회 발화. 마지막 내담자 반응은 \"{last_client}\". {status}."
|
||
)
|
||
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:
|
||
return _compact(text, limit=_PINNED_FACT_MAX_VALUE_CHARS)
|
||
|
||
|
||
def extract_pinned_fact_candidates(
|
||
masked_turns: list[dict[str, Any]],
|
||
) -> list[PinnedFactCandidate]:
|
||
"""Extract conservative pinned facts from masked client-visible text.
|
||
|
||
The first pass intentionally avoids clinical inference. It only preserves
|
||
explicit facts already surfaced by the client AI and already masked for
|
||
learner visibility.
|
||
"""
|
||
by_key: dict[str, PinnedFactCandidate] = {}
|
||
for turn in masked_turns:
|
||
if turn.get("speaker") != "client":
|
||
continue
|
||
text = _fact_value(turn.get("text"))
|
||
if not text:
|
||
continue
|
||
source_turn_id = turn.get("turn_id")
|
||
if "[NAME]" in text:
|
||
by_key["identity:name"] = PinnedFactCandidate(
|
||
key="identity:name",
|
||
value="[NAME]",
|
||
fact_type="identity",
|
||
confidence=0.85,
|
||
source_turn_id=str(source_turn_id) if source_turn_id else None,
|
||
)
|
||
if "[ORG]" in text:
|
||
by_key["identity:org"] = PinnedFactCandidate(
|
||
key="identity:org",
|
||
value="[ORG]",
|
||
fact_type="identity",
|
||
confidence=0.85,
|
||
source_turn_id=str(source_turn_id) if source_turn_id else None,
|
||
)
|
||
if _COUNSELING_AGREEMENT_WITHDRAWAL_RE.search(text):
|
||
by_key["agreement:counseling"] = PinnedFactCandidate(
|
||
key="agreement:counseling",
|
||
value=text,
|
||
fact_type="agreement",
|
||
status="contradicted",
|
||
confidence=0.8,
|
||
source_turn_id=str(source_turn_id) if source_turn_id else None,
|
||
)
|
||
elif _COUNSELING_AGREEMENT_RE.search(text):
|
||
by_key["agreement:counseling"] = PinnedFactCandidate(
|
||
key="agreement:counseling",
|
||
value=text,
|
||
fact_type="agreement",
|
||
confidence=0.75,
|
||
source_turn_id=str(source_turn_id) if source_turn_id else None,
|
||
)
|
||
return list(by_key.values())
|
||
|
||
|
||
def merge_case_digest(
|
||
*,
|
||
existing_digest: str | None,
|
||
session_no: int,
|
||
session_digest: str,
|
||
max_entries: int = _CASE_DIGEST_MAX_ENTRIES,
|
||
) -> str:
|
||
"""case_profile.case_digest를 session_no 기준으로 idempotent append한다."""
|
||
prefix = f"S{session_no}:"
|
||
lines = [
|
||
line.strip()
|
||
for line in str(existing_digest or "").splitlines()
|
||
if line.strip() and not line.strip().startswith(prefix)
|
||
]
|
||
next_line = session_digest.strip()
|
||
if next_line and not next_line.startswith(prefix):
|
||
next_line = f"{prefix} {next_line}"
|
||
if next_line:
|
||
lines.append(next_line)
|
||
return "\n".join(lines[-max_entries:])
|
||
|
||
|
||
def rapport_trajectory_point(*, session_no: int, end_state: dict) -> dict[str, Any]:
|
||
"""case_profile.rapport_trajectory에 저장할 최소 무손실 수치 포인트."""
|
||
return {
|
||
"session_no": int(session_no),
|
||
"stage": end_state.get("stage"),
|
||
"end_rapport": end_state.get("rapport_credit"),
|
||
"end_openness": end_state.get("effective_openness"),
|
||
"resistance": end_state.get("resistance"),
|
||
}
|
||
|
||
|
||
def merge_rapport_trajectory(
|
||
existing: Any,
|
||
point: dict[str, Any],
|
||
*,
|
||
max_entries: int = _RAPPORT_TRAJECTORY_MAX_ENTRIES,
|
||
) -> list[dict[str, Any]]:
|
||
"""session_no 기준으로 trajectory를 덮어쓰기 가능하게 append한다."""
|
||
session_no = point.get("session_no")
|
||
merged: list[dict[str, Any]] = []
|
||
if isinstance(existing, list):
|
||
for item in existing:
|
||
if not isinstance(item, dict):
|
||
continue
|
||
if item.get("session_no") == session_no:
|
||
continue
|
||
merged.append(dict(item))
|
||
merged.append(dict(point))
|
||
return merged[-max_entries:]
|
||
|
||
|
||
def update_alliance_level(previous: Any, end_rapport: Any) -> float:
|
||
"""case_profile.alliance_level EWMA. 이전 값이 없으면 schema default 0.2 기준."""
|
||
try:
|
||
prev = float(previous)
|
||
except (TypeError, ValueError):
|
||
prev = 0.2
|
||
try:
|
||
rapport = float(end_rapport)
|
||
except (TypeError, ValueError):
|
||
rapport = prev
|
||
return round(max(0.0, min(1.0, prev * 0.7 + rapport * 0.3)), 4)
|
||
|
||
|
||
__all__ = [
|
||
"RecallContext",
|
||
"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",
|
||
"merge_case_digest",
|
||
"rapport_trajectory_point",
|
||
"merge_rapport_trajectory",
|
||
"update_alliance_level",
|
||
]
|