vignette/apps/api/app/test_session_memory.py

950 lines
39 KiB
Python

from __future__ import annotations
import unittest
from contextlib import asynccontextmanager
from unittest.mock import patch
from . import session_persistence
from .routes import sessions
from .services import memory, rag, state_machine
from .services.persona import P1
from .store import InProcSession, TurnRecord
class SessionMemoryPureTest(unittest.TestCase):
def test_case_digest_merge_is_idempotent_by_session_no(self) -> None:
digest = memory.merge_case_digest(
existing_digest="S1: 이전 회기 요약\nS2: 오래된 요약",
session_no=2,
session_digest="S2: 새 요약",
)
self.assertEqual(digest, "S1: 이전 회기 요약\nS2: 새 요약")
def test_rapport_trajectory_merge_replaces_same_session(self) -> None:
merged = memory.merge_rapport_trajectory(
[{"session_no": 1, "end_rapport": 0.2}, {"session_no": 2, "end_rapport": 0.3}],
{"session_no": 2, "end_rapport": 0.7, "end_openness": 0.5},
)
self.assertEqual(len(merged), 2)
self.assertEqual(merged[-1]["session_no"], 2)
self.assertEqual(merged[-1]["end_rapport"], 0.7)
def test_fallback_session_digest_uses_masked_client_visible_text(self) -> None:
digest = memory.build_fallback_session_digest(
session_no=3,
masked_turns=[
{"speaker": "counselor", "text": "그때 마음이 어땠나요?"},
{"speaker": "client", "text": "저는 [NAME]이고 [ORG]에 다녀요."},
],
end_state={"stage": "탐색", "effective_openness": 0.42, "rapport_credit": 0.31},
)
self.assertIn("S3:", digest)
self.assertIn("[NAME]", digest)
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_llm_digest_worker_outcome_accepts_masked_contract_result(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=7,
masked_turns=[
{"speaker": "client", "text": "저는 [NAME]이고 가족 이야기가 어렵습니다.", "visible_to": ["client"]},
{"speaker": "counselor", "text": "그 주제를 다음 회기에 이어가겠습니다.", "visible_to": ["client"]},
],
open_threads=["가족 갈등을 다음 회기에 이어가기"],
)
outcome = memory.build_llm_digest_worker_outcome(
digest_input,
(
"내담자는 [NAME]으로 지칭되며 가족 갈등을 조심스럽게 설명했다. "
"상담자는 감정 확인과 다음 회기에서 이어갈 주제를 함께 정리했다."
),
forbidden_substrings=("김서연",),
)
self.assertFalse(outcome.fallback_required)
self.assertTrue(outcome.quality.accepted)
self.assertEqual(outcome.quality.reason, "ok")
self.assertIsNotNone(outcome.result)
assert outcome.result is not None
self.assertEqual(outcome.result.source, "llm")
self.assertEqual(outcome.result.session_no, 7)
self.assertEqual(outcome.result.open_threads, ("가족 갈등을 다음 회기에 이어가기",))
self.assertTrue(outcome.result.digest.startswith("S7:"))
self.assertIn("[NAME]", outcome.result.digest)
self.assertNotIn("김서연", outcome.result.digest)
def test_llm_digest_quality_rejects_raw_forbidden_substring(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=8,
masked_turns=[
{"speaker": "client", "text": "저는 [NAME]입니다.", "visible_to": ["client"]},
],
)
outcome = memory.build_llm_digest_worker_outcome(
digest_input,
"S8: 내담자 김서연은 가족 갈등을 설명했고 상담자는 다음 회기에서 이어갈 주제를 정리했다.",
forbidden_substrings=("김서연",),
)
self.assertTrue(outcome.fallback_required)
self.assertIsNone(outcome.result)
self.assertFalse(outcome.quality.accepted)
self.assertEqual(outcome.quality.reason, "forbidden_substring")
self.assertEqual(outcome.quality.details, ("김서연",))
def test_llm_digest_quality_rejects_internal_markers_and_wrong_session(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=9,
masked_turns=[
{"speaker": "client", "text": "저는 [NAME]입니다.", "visible_to": ["client"]},
],
)
marker_outcome = memory.build_llm_digest_worker_outcome(
digest_input,
"S9: rapport_credit 수치와 evaluation payload를 근거로 요약을 작성했다. 다음 회기 주제를 유지한다.",
)
wrong_session_outcome = memory.build_llm_digest_worker_outcome(
digest_input,
"S8: 내담자는 [NAME]으로 지칭되며 가족 갈등을 설명했다. 상담자는 다음 회기 주제를 정리했다.",
)
self.assertTrue(marker_outcome.fallback_required)
self.assertEqual(marker_outcome.quality.reason, "internal_marker")
self.assertTrue(wrong_session_outcome.fallback_required)
self.assertEqual(wrong_session_outcome.quality.reason, "wrong_session_prefix")
def test_extract_pinned_fact_candidates_is_conservative_and_masked(self) -> None:
facts = memory.extract_pinned_fact_candidates(
[
{"speaker": "counselor", "text": "이름을 말해줄 수 있나요?"},
{
"speaker": "client",
"text": "저는 [NAME]이고 [ORG]에 다녀요. 동생과 자주 싸워요. 주 1회 상담 약속은 지키고 싶어요.",
"turn_id": "00000000-0000-0000-0000-000000000201",
},
]
)
by_key = {fact.key: fact for fact in facts}
self.assertEqual(by_key["identity:name"].value, "[NAME]")
self.assertEqual(by_key["identity:org"].value, "[ORG]")
self.assertEqual(by_key["agreement:counseling"].fact_type, "agreement")
self.assertNotIn("relationship:sibling", by_key)
self.assertNotIn("김서연", " ".join(fact.value for fact in facts))
def test_extract_pinned_fact_candidates_skips_inferred_or_transient_content(self) -> None:
facts = memory.extract_pinned_fact_candidates(
[
{"speaker": "client", "text": "오늘은 그냥 기분이 좀 나빴어요."},
{"speaker": "client", "text": "동생과 자주 싸워요."},
{"speaker": "client", "text": "죽고 싶다는 생각이 스쳐갔어요."},
]
)
self.assertEqual(facts, [])
def test_extract_pinned_fact_candidates_marks_explicit_agreement_withdrawal_only(self) -> None:
facts = memory.extract_pinned_fact_candidates(
[
{
"speaker": "client",
"text": "상담 약속은 이제 못 지키겠어요. 회기는 그만하고 싶어요.",
"turn_id": "00000000-0000-0000-0000-000000000203",
},
{"speaker": "client", "text": "동생과 자주 싸워요."},
]
)
by_key = {fact.key: fact for fact in facts}
self.assertEqual(list(by_key), ["agreement:counseling"])
self.assertEqual(by_key["agreement:counseling"].status, "contradicted")
self.assertEqual(by_key["agreement:counseling"].fact_type, "agreement")
self.assertIn("못 지키겠어요", by_key["agreement:counseling"].value)
self.assertEqual(
by_key["agreement:counseling"].source_turn_id,
"00000000-0000-0000-0000-000000000203",
)
def test_episodic_turn_inputs_use_masked_client_visible_client_turns_only(self) -> None:
inputs = rag.episodic_turn_inputs_from_records(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
turns=[
TurnRecord(
turn_seq=1,
speaker="client",
stage="라포",
text="저는 김서연이고 한신대학교에 다녀요.",
text_masked="저는 [NAME]이고 [ORG]에 다녀요.",
turn_id="00000000-0000-0000-0000-000000000201",
),
TurnRecord(
turn_seq=2,
speaker="counselor",
stage="라포",
text="상담자 발화",
text_masked="상담자 발화",
turn_id="00000000-0000-0000-0000-000000000202",
),
TurnRecord(
turn_seq=3,
speaker="client",
stage="라포",
text="평가자만 볼 발화",
text_masked="평가자만 볼 발화",
turn_id="00000000-0000-0000-0000-000000000203",
visible_to=("evaluator",),
),
TurnRecord(
turn_seq=4,
speaker="client",
stage="라포",
text="DB turn id 없음",
text_masked="DB turn id 없음",
),
],
)
self.assertEqual(len(inputs), 1)
self.assertEqual(inputs[0].turn_id, "00000000-0000-0000-0000-000000000201")
self.assertEqual(inputs[0].seq, 1)
self.assertEqual(inputs[0].text_masked, "저는 [NAME]이고 [ORG]에 다녀요.")
self.assertNotIn("김서연", inputs[0].text_masked)
self.assertNotIn("한신대학교", inputs[0].text_masked)
class SessionMemoryPersistenceTest(unittest.IsolatedAsyncioTestCase):
async def test_write_persona_turn_embeddings_is_masked_and_idempotent(self) -> None:
class FakeConn:
def __init__(self) -> None:
self.executed: list[tuple[str, tuple[object, ...]]] = []
async def execute(self, query: str, *args: object) -> str:
self.executed.append((query, args))
return "INSERT 0 1"
conn = FakeConn()
turn = rag.EpisodicTurnInput(
turn_id="00000000-0000-0000-0000-000000000201",
case_id="00000000-0000-0000-0000-00000000ca5e",
session_id="00000000-0000-0000-0000-00000000feed",
seq=7,
text_masked="저는 [NAME]이고 [ORG]에 다녀요.",
)
captured_texts: list[str] = []
def fake_embed_query(text: str) -> rag.EmbeddedQuery:
captured_texts.append(text)
return rag.EmbeddedQuery(dense=[0.1] * rag.EMBED_DIM, sparse={"42": 0.7})
with patch.object(rag, "embed_query", fake_embed_query):
result = await rag.write_persona_turn_embeddings(conn, turns=[turn])
self.assertEqual(result.inserted, 1)
self.assertEqual(captured_texts, ["저는 [NAME]이고 [ORG]에 다녀요."])
self.assertEqual(len(conn.executed), 1)
query, args = conn.executed[0]
self.assertIn("INSERT INTO app.turn_embedding", query)
self.assertIn("ON CONFLICT (turn_id) DO NOTHING", query)
self.assertEqual(args[0], turn.turn_id)
self.assertEqual(args[1], turn.case_id)
self.assertEqual(args[2], turn.session_id)
self.assertEqual(args[3], 7)
self.assertIn("[0.1,0.1", str(args[4]))
self.assertEqual(args[5], '{"42": 0.7}')
self.assertNotIn("김서연", " ".join(str(arg) for arg in args))
async def test_end_persisted_session_schedules_episodic_embedding_writer(self) -> None:
scheduled: list[object] = []
def fake_create_task(coro):
scheduled.append(coro)
coro.close()
return None
sess = InProcSession(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
learner_id="00000000-0000-0000-0000-000000000101",
persona_code=P1.code,
theory_mode="humanistic",
persona=P1,
state=state_machine.SessionState(),
session_no=1,
)
carry = memory.CarryOver(
end_state={},
rapport_delta=0.0,
compression_job=None,
)
with patch.object(session_persistence, "end_session", return_value=True), patch.object(
sessions.asyncio,
"create_task",
fake_create_task,
):
await sessions._end_persisted_session(sess, carry)
self.assertTrue(sess.ended)
self.assertEqual(
[coro.cr_code.co_name for coro in scheduled],
["_write_episodic_embeddings", "close_session"],
)
async def test_end_persisted_session_schedules_digest_worker_only_when_enabled(self) -> None:
scheduled: list[str] = []
def fake_create_task(coro):
scheduled.append(coro.cr_code.co_name)
coro.close()
return None
sess = InProcSession(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
learner_id="00000000-0000-0000-0000-000000000101",
persona_code=P1.code,
theory_mode="humanistic",
persona=P1,
state=state_machine.SessionState(),
session_no=1,
)
digest_input = memory.build_session_digest_input(
session_id=sess.session_id,
case_id=sess.case_id,
session_no=sess.session_no,
masked_turns=[
{
"speaker": "client",
"text_masked": "다음 회기에 가족 이야기를 이어가고 싶어요.",
"visible_to": ["client"],
}
],
open_threads=["가족 이야기"],
)
carry = memory.CarryOver(
end_state={},
rapport_delta=0.0,
compression_job=memory.CompressionJob(digest_input=digest_input),
)
with patch.object(session_persistence, "end_session", return_value=True), patch.object(
sessions.settings,
"session_digest_worker_enabled",
True,
), patch.object(
sessions.asyncio,
"create_task",
fake_create_task,
):
await sessions._end_persisted_session(sess, carry)
self.assertEqual(
scheduled,
[
"_run_session_digest_worker_for_session",
"_write_episodic_embeddings",
"close_session",
],
)
async def test_end_persisted_session_keeps_digest_worker_default_off(self) -> None:
scheduled: list[str] = []
def fake_create_task(coro):
scheduled.append(coro.cr_code.co_name)
coro.close()
return None
sess = InProcSession(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
learner_id="00000000-0000-0000-0000-000000000101",
persona_code=P1.code,
theory_mode="humanistic",
persona=P1,
state=state_machine.SessionState(),
session_no=1,
)
digest_input = memory.build_session_digest_input(
session_id=sess.session_id,
case_id=sess.case_id,
session_no=sess.session_no,
masked_turns=[
{
"speaker": "client",
"text_masked": "다음 회기에 가족 이야기를 이어가고 싶어요.",
"visible_to": ["client"],
}
],
open_threads=["가족 이야기"],
)
carry = memory.CarryOver(
end_state={},
rapport_delta=0.0,
compression_job=memory.CompressionJob(digest_input=digest_input),
)
with patch.object(session_persistence, "end_session", return_value=True), patch.object(
sessions.settings,
"session_digest_worker_enabled",
False,
), patch.object(
sessions.asyncio,
"create_task",
fake_create_task,
):
await sessions._end_persisted_session(sess, carry)
self.assertEqual(scheduled, ["_write_episodic_embeddings", "close_session"])
async def test_session_digest_worker_releases_db_connection_during_engine_call(self) -> None:
order: list[str] = []
digest_input = memory.build_session_digest_input(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
session_no=1,
masked_turns=[
{
"speaker": "client",
"text_masked": "다음 회기에 가족 이야기를 이어가고 싶어요.",
"visible_to": ["client"],
}
],
open_threads=["가족 이야기"],
)
loaded = sessions.session_digest_worker.LoadedSessionDigestJob(
job=memory.CompressionJob(digest_input=digest_input),
existing_case_digest="S0: 이전",
learner_id="00000000-0000-0000-0000-000000000101",
)
class Worker:
apply_plan = object()
acquire_count = 0
@asynccontextmanager
async def fake_acquire(**kwargs):
nonlocal acquire_count
acquire_count += 1
label = "load" if acquire_count == 1 else "apply"
order.append(f"enter-{label}")
try:
yield object()
finally:
order.append(f"exit-{label}")
async def fake_load(conn, session_id: str):
order.append("load")
return loaded
async def fake_run(job, engine, *, existing_case_digest=None, model=None, audit_hook=None):
order.append("engine")
self.assertIs(engine, sessions.engine_client)
self.assertEqual(existing_case_digest, "S0: 이전")
self.assertIs(audit_hook, session_persistence.record_llm_call_audit)
return Worker()
async def fake_apply(conn, apply_plan, *, learner_id=None):
order.append("apply")
self.assertEqual(learner_id, loaded.learner_id)
return True
with patch.object(sessions.db, "get_pool", return_value=object()), patch.object(
sessions.db,
"acquire",
fake_acquire,
), patch.object(
sessions.session_digest_worker,
"load_session_digest_job",
fake_load,
), patch.object(
sessions.session_digest_worker,
"run_session_digest_worker",
fake_run,
), patch.object(
sessions.session_digest_worker,
"apply_session_digest_plan",
fake_apply,
), patch.object(
sessions.settings,
"session_digest_worker_model",
"",
):
await sessions._run_session_digest_worker_for_session(loaded.job.session_id)
self.assertEqual(
order,
["enter-load", "load", "exit-load", "engine", "enter-apply", "apply", "exit-apply"],
)
async def test_seed_recall_loads_case_digest_and_client_visible_pinned_facts(self) -> None:
test_case = self
class FakeConn:
def __init__(self) -> None:
self.fetch_queries: list[tuple[str, tuple[object, ...]]] = []
async def fetchrow(self, query: str, *args: object):
self.fetch_queries.append((query, args))
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.5},
}
return None
async def fetch(self, query: str, *args: object):
self.fetch_queries.append((query, args))
test_case.assertIn("FROM app.pinned_fact", query)
test_case.assertIn("status IN ('stable', 'evolving', 'locked')", query)
test_case.assertIn("$2 = ANY(visible_to)", query)
return [{"value": "동생과의 갈등"}, {"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
conn = FakeConn()
with patch.object(sessions.db, "get_pool", return_value=object()), patch.object(
sessions.db,
"acquire",
return_value=FakeAcquire(conn),
):
recall = await sessions._build_seed_recall(
case_id="00000000-0000-0000-0000-00000000ca5e"
)
self.assertIn("[케이스 큰그림]", recall.recall_summary or "")
self.assertIn("S1: 케이스 큰그림", recall.recall_summary or "")
self.assertIn("직전 회기 요약", recall.recall_summary or "")
self.assertEqual(recall.pinned_facts, ["동생과의 갈등", "주 1회 상담 약속"])
self.assertEqual(recall.carry, {"rapport_credit": 0.5})
async def test_end_session_updates_session_summary_and_case_profile(self) -> None:
test_case = self
class FakeConn:
def __init__(self) -> None:
self.executed: list[tuple[str, tuple[object, ...]]] = []
async def execute(self, query: str, *args: object) -> str:
self.executed.append((query, args))
return "OK"
async def fetchrow(self, query: str, *args: object):
self.executed.append((query, args))
if "INSERT INTO app.pinned_fact" in query:
return {
"id": f"00000000-0000-0000-0000-00000000fa{len(self.executed):02d}",
"old_value": None,
"new_value": args[2],
}
test_case.assertIn("FROM app.case_profile", query)
return {
"case_digest": "S1: 이전 회기",
"rapport_trajectory": [{"session_no": 1, "end_rapport": 0.2}],
"alliance_level": 0.2,
}
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
learner_id = "00000000-0000-0000-0000-000000000101"
case_id = "00000000-0000-0000-0000-00000000ca5e"
state = state_machine.SessionState(
stage=state_machine.Stage.EXPLORE,
turn_seq=2,
effective_openness=0.42,
rapport_credit=0.31,
resistance=0.5,
)
sess = InProcSession(
session_id="00000000-0000-0000-0000-00000000feed",
case_id=case_id,
learner_id=learner_id,
persona_code=P1.code,
theory_mode="humanistic",
persona=P1,
state=state,
session_no=2,
prev_rapport_credit=0.1,
turns=[
TurnRecord(
turn_seq=1,
speaker="counselor",
stage="라포",
text="실명 질문",
text_masked="이름을 말해줄 수 있나요?",
),
TurnRecord(
turn_seq=1,
speaker="client",
stage="라포",
text="저는 김서연이고 한신대학교에 다녀요. 동생과 자주 싸워요. 주 1회 상담 약속은 지키고 싶어요.",
text_masked="저는 [NAME]이고 [ORG]에 다녀요. 동생과 자주 싸워요. 주 1회 상담 약속은 지키고 싶어요.",
turn_id="00000000-0000-0000-0000-000000000201",
),
],
)
carry = memory.make_carry_over(
state=state,
session_id=sess.session_id,
case_id=sess.case_id,
session_no=sess.session_no,
masked_turns=sess.masked_turns(),
prev_rapport_credit=sess.prev_rapport_credit,
)
conn = FakeConn()
with patch.object(session_persistence, "get_pool", return_value=object()), patch.object(
session_persistence,
"acquire",
return_value=FakeAcquire(conn),
):
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(summary_digest, case_digest)
self.assertIn("[NAME]", case_digest)
self.assertNotIn("김서연", case_digest)
self.assertEqual(trajectory[-1]["session_no"], 2)
self.assertEqual(trajectory[-1]["end_rapport"], 0.31)
self.assertGreater(alliance_level, 0.2)
pinned_writes = [
args
for query, args in conn.executed
if "INSERT INTO app.pinned_fact (" in query
]
self.assertEqual(len(pinned_writes), 3)
by_key = {args[1]: args for args in pinned_writes}
self.assertEqual(by_key["identity:name"][2], "[NAME]")
self.assertEqual(by_key["identity:org"][2], "[ORG]")
self.assertEqual(by_key["agreement:counseling"][3], "agreement")
self.assertNotIn("relationship:sibling", by_key)
self.assertEqual(by_key["identity:name"][5], "00000000-0000-0000-0000-000000000201")
self.assertEqual(by_key["identity:name"][7], 2)
self.assertEqual(by_key["identity:name"][8], ["client", "evaluator"])
self.assertNotIn("김서연", " ".join(str(args[2]) for args in pinned_writes))
history_writes = [
args for query, args in conn.executed if "INSERT INTO app.pinned_fact_history" in query
]
self.assertEqual(len(history_writes), 3)
self.assertEqual(history_writes[0][1], case_id)
self.assertIsNone(history_writes[0][2])
self.assertEqual(history_writes[0][3], "[NAME]")
self.assertEqual(history_writes[0][4], "progression")
self.assertEqual(history_writes[0][5], 2)
self.assertEqual(
history_writes[0][6],
"00000000-0000-0000-0000-000000000201",
)
self.assertNotIn("김서연", " ".join(str(args[3]) for args in history_writes))
async def test_pinned_fact_history_skips_same_value_refresh(self) -> None:
class FakeConn:
def __init__(self) -> None:
self.executed: list[tuple[str, tuple[object, ...]]] = []
async def fetchrow(self, query: str, *args: object):
self.executed.append((query, args))
return {
"id": "00000000-0000-0000-0000-00000000fa11",
"old_value": args[2],
"new_value": args[2],
}
async def execute(self, query: str, *args: object) -> str:
self.executed.append((query, args))
return "OK"
conn = FakeConn()
sess = InProcSession(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
learner_id="00000000-0000-0000-0000-000000000101",
persona_code=P1.code,
theory_mode="humanistic",
persona=P1,
state=state_machine.SessionState(),
session_no=3,
turns=[
TurnRecord(
turn_seq=1,
speaker="client",
stage="라포",
text="저는 김서연입니다.",
text_masked="저는 [NAME]입니다.",
turn_id="00000000-0000-0000-0000-000000000201",
)
],
)
await session_persistence._upsert_pinned_fact_candidates(conn, sess)
pinned_writes = [
args
for query, args in conn.executed
if "INSERT INTO app.pinned_fact (" in query
]
history_writes = [
args for query, args in conn.executed if "INSERT INTO app.pinned_fact_history" in query
]
self.assertEqual(len(pinned_writes), 1)
self.assertEqual(history_writes, [])
async def test_pinned_fact_history_records_value_change_as_clarification(self) -> None:
class FakeConn:
def __init__(self) -> None:
self.executed: list[tuple[str, tuple[object, ...]]] = []
async def fetchrow(self, query: str, *args: object):
self.executed.append((query, args))
return {
"id": "00000000-0000-0000-0000-00000000fa22",
"old_value": "예전에는 격주 상담 약속을 말함.",
"new_value": args[2],
}
async def execute(self, query: str, *args: object) -> str:
self.executed.append((query, args))
return "OK"
conn = FakeConn()
sess = InProcSession(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
learner_id="00000000-0000-0000-0000-000000000101",
persona_code=P1.code,
theory_mode="humanistic",
persona=P1,
state=state_machine.SessionState(),
session_no=4,
turns=[
TurnRecord(
turn_seq=1,
speaker="client",
stage="라포",
text="앞으로는 주 1회 상담 약속을 지키고 싶어요.",
text_masked="앞으로는 주 1회 상담 약속을 지키고 싶어요.",
turn_id="00000000-0000-0000-0000-000000000202",
)
],
)
await session_persistence._upsert_pinned_fact_candidates(conn, sess)
history_writes = [
args for query, args in conn.executed if "INSERT INTO app.pinned_fact_history" in query
]
self.assertEqual(len(history_writes), 1)
self.assertEqual(history_writes[0][2], "예전에는 격주 상담 약속을 말함.")
self.assertIn("주 1회 상담 약속", str(history_writes[0][3]))
self.assertEqual(history_writes[0][4], "clarification")
self.assertEqual(history_writes[0][5], 4)
async def test_pinned_fact_history_records_explicit_agreement_withdrawal_as_contradiction(self) -> None:
test_case = self
class FakeConn:
def __init__(self) -> None:
self.executed: list[tuple[str, tuple[object, ...]]] = []
async def fetchrow(self, query: str, *args: object):
self.executed.append((query, args))
test_case.assertIn("UPDATE app.pinned_fact", query)
test_case.assertIn("status = 'contradicted'", query)
return {
"id": "00000000-0000-0000-0000-00000000fa33",
"old_value": "주 1회 상담 약속은 지키고 싶어요.",
"new_value": args[2],
}
async def execute(self, query: str, *args: object) -> str:
self.executed.append((query, args))
return "OK"
conn = FakeConn()
sess = InProcSession(
session_id="00000000-0000-0000-0000-00000000feed",
case_id="00000000-0000-0000-0000-00000000ca5e",
learner_id="00000000-0000-0000-0000-000000000101",
persona_code=P1.code,
theory_mode="humanistic",
persona=P1,
state=state_machine.SessionState(),
session_no=5,
turns=[
TurnRecord(
turn_seq=1,
speaker="client",
stage="라포",
text="상담 약속은 이제 못 지키겠어요. 회기는 그만하고 싶어요.",
text_masked="상담 약속은 이제 못 지키겠어요. 회기는 그만하고 싶어요.",
turn_id="00000000-0000-0000-0000-000000000203",
)
],
)
await session_persistence._upsert_pinned_fact_candidates(conn, sess)
contradiction_updates = [
args for query, args in conn.executed if "UPDATE app.pinned_fact" in query
]
self.assertEqual(len(contradiction_updates), 1)
self.assertEqual(contradiction_updates[0][1], "agreement:counseling")
self.assertIn("그만하고 싶어요", str(contradiction_updates[0][2]))
self.assertEqual(contradiction_updates[0][7], ["evaluator"])
history_writes = [
args for query, args in conn.executed if "INSERT INTO app.pinned_fact_history" in query
]
self.assertEqual(len(history_writes), 1)
self.assertEqual(history_writes[0][2], "주 1회 상담 약속은 지키고 싶어요.")
self.assertIn("못 지키겠어요", str(history_writes[0][3]))
self.assertEqual(history_writes[0][4], "contradiction")
self.assertEqual(history_writes[0][5], 5)
if __name__ == "__main__":
unittest.main()