vignette/apps/api/app/services/memory.py
2026-07-15 21:31:30 +09:00

676 lines
24 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""회기 라이프사이클 메모리 — 시작 회상 + 종료 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, Literal, Optional
from .state_machine import SessionState
from ..taxonomy import speaker_ko_label
_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"}
_LLM_DIGEST_MIN_CHARS = 60
_LLM_DIGEST_INTERNAL_MARKERS = (
"rapport_credit",
"effective_openness",
"end_state",
"evaluation",
"정답",
"점수",
"평가 payload",
"평가점수",
"평가 점수",
"core_belief",
"CCD",
)
_SESSION_DIGEST_PREFIX_RE = re.compile(r"^S(?P<session_no>\d+):")
# ════════════════════════════════════════════════════════════════════════════
# 회기 시작 — 회상 (큰그림 → 세부)
# ════════════════════════════════════════════════════════════════════════════
@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 LLM worker paths."""
session_id: str
case_id: str | None
session_no: int
digest: str
open_threads: tuple[str, ...] = ()
source: Literal["fallback", "llm"] = "fallback"
@dataclass(frozen=True, slots=True)
class DigestQualityAssessment:
"""Local quality gate result before an LLM digest can replace fallback."""
accepted: bool
reason: Literal[
"ok",
"empty",
"too_short",
"forbidden_substring",
"internal_marker",
"wrong_session_prefix",
]
retryable: bool = False
details: tuple[str, ...] = ()
@dataclass(frozen=True, slots=True)
class SessionDigestWorkerOutcome:
"""Accepted LLM digest result or the reason fallback must remain authoritative."""
result: SessionDigestResult | None
quality: DigestQualityAssessment
@property
def fallback_required(self) -> bool:
return self.result is None
@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"{speaker_ko_label(t.speaker)}: {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 _normalize_digest_text(value: Any) -> str:
return " ".join(str(value or "").split())
def assess_llm_digest_quality(
digest_input: SessionDigestInput,
digest: Any,
*,
forbidden_substrings: tuple[str, ...] = (),
min_chars: int = _LLM_DIGEST_MIN_CHARS,
) -> DigestQualityAssessment:
"""Validate an LLM digest candidate without using raw text or end-state data."""
text = _normalize_digest_text(digest)
if not text:
return DigestQualityAssessment(accepted=False, reason="empty", retryable=True)
prefix_match = _SESSION_DIGEST_PREFIX_RE.match(text)
if prefix_match and int(prefix_match.group("session_no")) != digest_input.session_no:
return DigestQualityAssessment(
accepted=False,
reason="wrong_session_prefix",
retryable=True,
details=(prefix_match.group(0),),
)
forbidden_hits = tuple(
marker
for marker in (str(item).strip() for item in forbidden_substrings)
if marker and marker in text
)
if forbidden_hits:
return DigestQualityAssessment(
accepted=False,
reason="forbidden_substring",
retryable=True,
details=forbidden_hits[:5],
)
lowered = text.lower()
marker_hits = tuple(
marker
for marker in _LLM_DIGEST_INTERNAL_MARKERS
if marker.lower() in lowered
)
if marker_hits:
return DigestQualityAssessment(
accepted=False,
reason="internal_marker",
retryable=True,
details=marker_hits,
)
if len(text) < max(1, int(min_chars)):
return DigestQualityAssessment(accepted=False, reason="too_short", retryable=True)
return DigestQualityAssessment(accepted=True, reason="ok")
def build_llm_digest_worker_outcome(
digest_input: SessionDigestInput,
digest: Any,
*,
forbidden_substrings: tuple[str, ...] = (),
min_chars: int = _LLM_DIGEST_MIN_CHARS,
) -> SessionDigestWorkerOutcome:
"""Coerce a candidate LLM digest into the shared result contract if it passes."""
quality = assess_llm_digest_quality(
digest_input,
digest,
forbidden_substrings=forbidden_substrings,
min_chars=min_chars,
)
if not quality.accepted:
return SessionDigestWorkerOutcome(result=None, quality=quality)
normalized = _normalize_digest_text(digest)
prefix = f"S{digest_input.session_no}:"
if not normalized.startswith(prefix):
normalized = f"{prefix} {normalized}"
return SessionDigestWorkerOutcome(
result=SessionDigestResult(
session_id=digest_input.session_id,
case_id=digest_input.case_id,
session_no=digest_input.session_no,
digest=normalized,
open_threads=digest_input.open_threads,
source="llm",
),
quality=quality,
)
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",
"DigestQualityAssessment",
"SessionDigestWorkerOutcome",
"make_carry_over",
"build_session_digest_input",
"build_compression_messages",
"assess_llm_digest_quality",
"build_llm_digest_worker_outcome",
"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",
]