312 lines
12 KiB
Python
312 lines
12 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
|
|
from .contracts.engine_gateway import GenerateResponse
|
|
from .services import memory, session_digest_worker
|
|
|
|
|
|
SESSION_ID = "00000000-0000-0000-0000-00000000feed"
|
|
CASE_ID = "00000000-0000-0000-0000-00000000ca5e"
|
|
LEARNER_ID = "00000000-0000-0000-0000-000000000101"
|
|
|
|
|
|
def _digest_input(session_no: int = 2) -> memory.SessionDigestInput:
|
|
return memory.build_session_digest_input(
|
|
session_id=SESSION_ID,
|
|
case_id=CASE_ID,
|
|
session_no=session_no,
|
|
masked_turns=[
|
|
{
|
|
"speaker": "client",
|
|
"text": "저는 [NAME]이고 가족 갈등을 조심스럽게 설명했습니다.",
|
|
"visible_to": ["client", "evaluator"],
|
|
},
|
|
{
|
|
"speaker": "counselor",
|
|
"text": "그 이야기는 다음 회기에서 천천히 이어가겠습니다.",
|
|
"visible_to": ["client"],
|
|
},
|
|
{
|
|
"speaker": "client",
|
|
"text": "평가자 전용 발화",
|
|
"visible_to": ["evaluator"],
|
|
},
|
|
],
|
|
open_threads=["가족 갈등을 다음 회기에서 이어가기"],
|
|
)
|
|
|
|
|
|
def _accepted_text(session_no: int = 2) -> str:
|
|
return (
|
|
f"S{session_no}: 내담자는 [NAME]으로 지칭되며 가족 갈등을 조심스럽게 꺼냈다. "
|
|
"상담자는 감정을 서두르지 않고 확인했고 다음 회기에서 같은 주제를 이어가기로 했다. "
|
|
"내담자는 관계 이야기를 계속 다루는 데 약간의 부담과 기대를 함께 보였다."
|
|
)
|
|
|
|
|
|
class FakeEngine:
|
|
def __init__(self, text: str) -> None:
|
|
self.text = text
|
|
self.requests = []
|
|
|
|
async def generate(self, req):
|
|
self.requests.append(req)
|
|
return GenerateResponse(
|
|
text=self.text,
|
|
provider="openai",
|
|
model="gpt-4.1-mini",
|
|
tokens_in=20,
|
|
tokens_out=13,
|
|
cost_usd=0.001,
|
|
inference_geo="us",
|
|
)
|
|
|
|
|
|
class SessionDigestWorkerPureTest(unittest.IsolatedAsyncioTestCase):
|
|
def test_request_uses_gateway_contract_without_internal_state(self) -> None:
|
|
job = memory.CompressionJob(digest_input=_digest_input())
|
|
|
|
request = session_digest_worker.build_session_digest_request(job)
|
|
prompt = "\n".join(message.content for message in request.messages)
|
|
|
|
self.assertEqual(request.ai_role, "evaluator")
|
|
self.assertEqual(request.session_id, SESSION_ID)
|
|
self.assertEqual(request.metadata["loop"], "session_digest")
|
|
self.assertEqual(request.metadata["case_id"], CASE_ID)
|
|
self.assertEqual(request.metadata["session_no"], 2)
|
|
self.assertEqual([message.role for message in request.messages], ["system", "user"])
|
|
self.assertIn("[NAME]", prompt)
|
|
self.assertIn("가족 갈등을 다음 회기에서 이어가기", prompt)
|
|
self.assertNotIn("김서연", prompt)
|
|
self.assertNotIn("end_state", prompt)
|
|
self.assertNotIn("rapport_credit", prompt)
|
|
self.assertNotIn("CCD", prompt)
|
|
self.assertNotIn("평가자 전용", prompt)
|
|
|
|
async def test_run_builds_apply_plan_and_audits_accepted_digest(self) -> None:
|
|
job = memory.CompressionJob(digest_input=_digest_input())
|
|
engine = FakeEngine(_accepted_text())
|
|
audits = []
|
|
|
|
async def audit_hook(payload):
|
|
audits.append(payload)
|
|
|
|
run = await session_digest_worker.run_session_digest_worker(
|
|
job,
|
|
engine,
|
|
existing_case_digest="S1: 이전 회기\nS2: 오래된 요약",
|
|
audit_hook=audit_hook,
|
|
)
|
|
|
|
self.assertFalse(run.fallback_required)
|
|
self.assertIsNotNone(run.apply_plan)
|
|
assert run.apply_plan is not None
|
|
self.assertEqual(run.apply_plan.compressed_by, "llm:openai/gpt-4.1-mini")
|
|
self.assertEqual(run.apply_plan.token_count, 33)
|
|
self.assertIn("S1: 이전 회기", run.apply_plan.case_digest or "")
|
|
self.assertIn(_accepted_text(), run.apply_plan.case_digest or "")
|
|
self.assertNotIn("오래된 요약", run.apply_plan.case_digest or "")
|
|
self.assertEqual(audits[0]["session_id"], SESSION_ID)
|
|
self.assertEqual(audits[0]["provider"], "openai")
|
|
self.assertEqual(audits[0]["tokens_in"], 20)
|
|
self.assertEqual(len(engine.requests), 1)
|
|
|
|
async def test_rejected_digest_does_not_create_apply_plan(self) -> None:
|
|
job = memory.CompressionJob(digest_input=_digest_input(session_no=3))
|
|
engine = FakeEngine("S2: 내담자 김서연은 잘못된 회기 prefix와 raw 이름을 포함했다.")
|
|
|
|
run = await session_digest_worker.run_session_digest_worker(
|
|
job,
|
|
engine,
|
|
forbidden_substrings=("김서연",),
|
|
)
|
|
|
|
self.assertTrue(run.fallback_required)
|
|
self.assertIsNone(run.apply_plan)
|
|
self.assertTrue(run.outcome.fallback_required)
|
|
self.assertIn(run.outcome.quality.reason, {"wrong_session_prefix", "forbidden_substring"})
|
|
|
|
|
|
class SessionDigestWorkerPersistenceTest(unittest.IsolatedAsyncioTestCase):
|
|
async def test_load_session_digest_job_rebuilds_masked_client_visible_input(self) -> None:
|
|
class FakeConn:
|
|
def __init__(self) -> None:
|
|
self.fetchrow_query = ""
|
|
|
|
async def fetchrow(self, query: str, *args):
|
|
self.fetchrow_query = query
|
|
return {
|
|
"session_id": SESSION_ID,
|
|
"case_id": CASE_ID,
|
|
"session_no": 4,
|
|
"open_threads": ["다음 회기 주제"],
|
|
"case_digest": "S3: 이전 회기",
|
|
"learner_id": LEARNER_ID,
|
|
}
|
|
|
|
async def fetch(self, query: str, *args):
|
|
return [
|
|
{
|
|
"id": "00000000-0000-0000-0000-000000000201",
|
|
"speaker": "client",
|
|
"text": "raw 김서연",
|
|
"text_masked": "저는 [NAME]입니다.",
|
|
"visible_to": ["client", "evaluator"],
|
|
},
|
|
{
|
|
"id": "00000000-0000-0000-0000-000000000202",
|
|
"speaker": "client",
|
|
"text": "평가자 전용 raw",
|
|
"text_masked": "평가자 전용 masked",
|
|
"visible_to": ["evaluator"],
|
|
},
|
|
{
|
|
"id": "00000000-0000-0000-0000-000000000203",
|
|
"speaker": "client",
|
|
"text": "raw 김서연",
|
|
"text_masked": "",
|
|
"visible_to": ["client"],
|
|
},
|
|
{
|
|
"id": "00000000-0000-0000-0000-000000000204",
|
|
"speaker": "system",
|
|
"text": "시스템 메모",
|
|
"text_masked": "시스템 메모",
|
|
"visible_to": ["client"],
|
|
},
|
|
]
|
|
|
|
conn = FakeConn()
|
|
loaded = await session_digest_worker.load_session_digest_job(conn, SESSION_ID)
|
|
|
|
self.assertIsNotNone(loaded)
|
|
assert loaded is not None
|
|
self.assertIn("ss.compressed_by IS NULL", conn.fetchrow_query)
|
|
digest_input = loaded.job.digest_input
|
|
self.assertEqual(digest_input.session_no, 4)
|
|
self.assertEqual(digest_input.open_threads, ("다음 회기 주제",))
|
|
self.assertEqual(len(digest_input.masked_turns), 1)
|
|
self.assertEqual(digest_input.masked_turns[0].text, "저는 [NAME]입니다.")
|
|
self.assertNotIn("김서연", " ".join(turn.text for turn in digest_input.masked_turns))
|
|
self.assertEqual(loaded.existing_case_digest, "S3: 이전 회기")
|
|
self.assertEqual(loaded.learner_id, LEARNER_ID)
|
|
|
|
async def test_load_session_digest_job_skips_already_compressed_rows(self) -> None:
|
|
class FakeConn:
|
|
def __init__(self) -> None:
|
|
self.fetchrow_query = ""
|
|
self.fetch_called = False
|
|
|
|
async def fetchrow(self, query: str, *args):
|
|
self.fetchrow_query = query
|
|
return None
|
|
|
|
async def fetch(self, query: str, *args):
|
|
self.fetch_called = True
|
|
return []
|
|
|
|
conn = FakeConn()
|
|
loaded = await session_digest_worker.load_session_digest_job(conn, SESSION_ID)
|
|
|
|
self.assertIsNone(loaded)
|
|
self.assertIn("ss.compressed_by IS NULL", conn.fetchrow_query)
|
|
self.assertFalse(conn.fetch_called)
|
|
|
|
async def test_run_session_digest_once_applies_only_accepted_plan(self) -> None:
|
|
class FakeConn:
|
|
def __init__(self) -> None:
|
|
self.executed = []
|
|
|
|
async def fetchrow(self, query: str, *args):
|
|
return {
|
|
"session_id": SESSION_ID,
|
|
"case_id": CASE_ID,
|
|
"session_no": 5,
|
|
"open_threads": ["가족 갈등"],
|
|
"case_digest": "S4: 이전 회기\nS5: 오래된 요약",
|
|
"learner_id": LEARNER_ID,
|
|
}
|
|
|
|
async def fetch(self, query: str, *args):
|
|
return [
|
|
{
|
|
"id": "00000000-0000-0000-0000-000000000201",
|
|
"speaker": "client",
|
|
"text": "raw 김서연",
|
|
"text_masked": "저는 [NAME]이고 가족 갈등이 부담됩니다.",
|
|
"visible_to": ["client"],
|
|
},
|
|
{
|
|
"id": "00000000-0000-0000-0000-000000000202",
|
|
"speaker": "counselor",
|
|
"text": "다음 회기에 이어가겠습니다.",
|
|
"text_masked": "다음 회기에 이어가겠습니다.",
|
|
"visible_to": ["client"],
|
|
},
|
|
]
|
|
|
|
async def execute(self, query: str, *args):
|
|
self.executed.append((query, args))
|
|
return "UPDATE 1"
|
|
|
|
conn = FakeConn()
|
|
result = await session_digest_worker.run_session_digest_once(
|
|
conn,
|
|
session_id=SESSION_ID,
|
|
engine=FakeEngine(_accepted_text(session_no=5)),
|
|
)
|
|
|
|
self.assertTrue(result.found)
|
|
self.assertTrue(result.applied)
|
|
self.assertEqual(len(conn.executed), 2)
|
|
summary_query, summary_args = conn.executed[0]
|
|
self.assertIn("UPDATE app.session_summary", summary_query)
|
|
self.assertEqual(summary_args[0], SESSION_ID)
|
|
self.assertIn("[NAME]", summary_args[1])
|
|
self.assertEqual(summary_args[2], ["가족 갈등"])
|
|
self.assertEqual(summary_args[3], "llm:openai/gpt-4.1-mini")
|
|
self.assertEqual(summary_args[4], 33)
|
|
case_query, case_args = conn.executed[1]
|
|
self.assertIn("UPDATE app.case_profile", case_query)
|
|
self.assertEqual(case_args[0], CASE_ID)
|
|
self.assertEqual(case_args[1], LEARNER_ID)
|
|
self.assertIn("S4: 이전 회기", case_args[2])
|
|
self.assertIn(_accepted_text(session_no=5), case_args[2])
|
|
self.assertNotIn("오래된 요약", case_args[2])
|
|
|
|
async def test_apply_session_digest_plan_uses_cas_before_case_update(self) -> None:
|
|
class FakeConn:
|
|
def __init__(self) -> None:
|
|
self.executed = []
|
|
|
|
async def execute(self, query: str, *args):
|
|
self.executed.append((query, args))
|
|
return "UPDATE 0"
|
|
|
|
plan = session_digest_worker.SessionDigestApplyPlan(
|
|
result=memory.SessionDigestResult(
|
|
session_id=SESSION_ID,
|
|
case_id=CASE_ID,
|
|
session_no=5,
|
|
digest=_accepted_text(session_no=5),
|
|
open_threads=("가족 갈등",),
|
|
source="llm",
|
|
),
|
|
case_digest="S5: accepted",
|
|
compressed_by="llm:openai/gpt-4.1-mini",
|
|
token_count=33,
|
|
)
|
|
conn = FakeConn()
|
|
|
|
applied = await session_digest_worker.apply_session_digest_plan(
|
|
conn,
|
|
plan,
|
|
learner_id=LEARNER_ID,
|
|
)
|
|
|
|
self.assertFalse(applied)
|
|
self.assertEqual(len(conn.executed), 1)
|
|
summary_query, summary_args = conn.executed[0]
|
|
self.assertIn("AND compressed_by IS NULL", summary_query)
|
|
self.assertEqual(summary_args[0], SESSION_ID)
|