323 lines
11 KiB
Python
323 lines
11 KiB
Python
"""Regression tests for P1 PII masking before engine requests."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import unittest
|
|
from typing import Any
|
|
from unittest.mock import patch
|
|
|
|
from .contracts.engine_gateway import EngineGatewaySseLineDecoder
|
|
from .engine_client import EngineClient, GenerateResponse
|
|
from .services import guardrail, orchestrator, persona, state_machine
|
|
|
|
|
|
RAW_PHONE = "010-1234-5678"
|
|
RAW_EMAIL = "test@example.com"
|
|
RAW_RRN = "990101-1234567"
|
|
RAW_TEXT = f"My phone is {RAW_PHONE}, email {RAW_EMAIL}, and RRN {RAW_RRN}."
|
|
RAW_VALUES = (RAW_PHONE, RAW_EMAIL, RAW_RRN)
|
|
MASK_VALUES = ("[PHONE]", "[EMAIL]", "[RRN]")
|
|
RAW_KO_NAME = "김서연"
|
|
RAW_KO_ORG = "한신대학교"
|
|
RAW_KO_DEPT = "상담심리학과"
|
|
RAW_KO_TEXT = (
|
|
f"내담자 {RAW_KO_NAME}은 {RAW_KO_ORG} {RAW_KO_DEPT} 학생이고 "
|
|
"연락은 하지 말아 주세요."
|
|
)
|
|
RAW_KO_VALUES = (RAW_KO_NAME, RAW_KO_ORG, RAW_KO_DEPT)
|
|
MASK_KO_VALUES = ("[NAME]", "[ORG]")
|
|
|
|
|
|
def _json_blob(value: Any) -> str:
|
|
return json.dumps(value, ensure_ascii=False, sort_keys=True, default=str)
|
|
|
|
|
|
def _message_blob(messages: object) -> str:
|
|
return "\n".join(message.content for message in messages) # type: ignore[attr-defined]
|
|
|
|
|
|
def _initial_state() -> state_machine.SessionState:
|
|
card = persona.P1
|
|
return state_machine.init_state(
|
|
params=card.openness_params(),
|
|
)
|
|
|
|
|
|
def _prepare_context() -> orchestrator.TurnContext:
|
|
return orchestrator.prepare_turn(
|
|
session_id="masking-session",
|
|
case_id="masking-case",
|
|
card=persona.P1,
|
|
state=_initial_state(),
|
|
learner_text=RAW_TEXT,
|
|
recent_turns=[
|
|
{
|
|
"speaker": "counselor",
|
|
"text": "Previous learner contact was already masked: [PHONE] [EMAIL] [RRN].",
|
|
}
|
|
],
|
|
)
|
|
|
|
|
|
def _assert_no_raw_pii(test: unittest.TestCase, value: object) -> None:
|
|
blob = _json_blob(value)
|
|
for raw in RAW_VALUES:
|
|
test.assertNotIn(raw, blob)
|
|
|
|
|
|
def _assert_masked_pii_present(test: unittest.TestCase, value: object) -> None:
|
|
blob = _json_blob(value)
|
|
for masked in MASK_VALUES:
|
|
test.assertIn(masked, blob)
|
|
|
|
|
|
def _assert_no_raw_ko_pii(test: unittest.TestCase, value: object) -> None:
|
|
blob = _json_blob(value)
|
|
for raw in RAW_KO_VALUES:
|
|
test.assertNotIn(raw, blob)
|
|
|
|
|
|
def _assert_masked_ko_pii_present(test: unittest.TestCase, value: object) -> None:
|
|
blob = _json_blob(value)
|
|
for masked in MASK_KO_VALUES:
|
|
test.assertIn(masked, blob)
|
|
|
|
|
|
class CaptureGenerateEngine:
|
|
def __init__(self) -> None:
|
|
self.request = None
|
|
self.payload: dict[str, Any] | None = None
|
|
self._payload_builder = EngineClient(base_url="http://engine.test")
|
|
|
|
async def generate(self, req):
|
|
self.request = req
|
|
self.payload = self._payload_builder._payload(req)
|
|
return GenerateResponse(
|
|
text="Masked engine reply.",
|
|
model="fake-model",
|
|
provider="fake-provider",
|
|
tokens_in=3,
|
|
tokens_out=4,
|
|
cost_usd=0.0,
|
|
)
|
|
|
|
|
|
class CaptureStreamEngine:
|
|
engine_mode = "fake-provider"
|
|
default_model = "fake-model"
|
|
|
|
def __init__(self) -> None:
|
|
self.request = None
|
|
self.payload: dict[str, Any] | None = None
|
|
self._payload_builder = EngineClient(base_url="http://engine.test")
|
|
|
|
async def stream(self, req):
|
|
self.request = req
|
|
self.payload = self._payload_builder._payload(req)
|
|
yield "event: token"
|
|
yield 'data: {"text":"Masked stream reply."}'
|
|
yield "event: done"
|
|
yield (
|
|
'data: {"provider":"fake-provider","model":"fake-model",'
|
|
'"tokens_in":5,"tokens_out":6,"cost_usd":0.0}'
|
|
)
|
|
|
|
async def stream_packets(self, req):
|
|
decoder = EngineGatewaySseLineDecoder()
|
|
async for raw in self.stream(req):
|
|
packet = decoder.feed_line(raw)
|
|
if packet is not None:
|
|
yield packet
|
|
|
|
|
|
class OrchestratorMaskingGateTest(unittest.IsolatedAsyncioTestCase):
|
|
def setUp(self) -> None:
|
|
self.presidio_patch = patch.object(
|
|
guardrail,
|
|
"_try_load_presidio",
|
|
return_value=(None, None),
|
|
)
|
|
self.presidio_patch.start()
|
|
self.addCleanup(self.presidio_patch.stop)
|
|
|
|
def test_prepare_turn_keeps_raw_text_but_builds_masked_engine_messages(self) -> None:
|
|
ctx = _prepare_context()
|
|
|
|
self.assertEqual(ctx.learner_text_raw, RAW_TEXT)
|
|
for raw in RAW_VALUES:
|
|
self.assertIn(raw, ctx.learner_text_raw)
|
|
self.assertNotIn(raw, ctx.learner_text_masked)
|
|
self.assertNotIn(raw, _message_blob(ctx.messages))
|
|
|
|
for masked in MASK_VALUES:
|
|
self.assertIn(masked, ctx.learner_text_masked)
|
|
self.assertIn(masked, _message_blob(ctx.messages))
|
|
|
|
def test_prepare_turn_masks_raw_pii_from_context_inputs(self) -> None:
|
|
ctx = orchestrator.prepare_turn(
|
|
session_id="masking-session",
|
|
case_id="masking-case",
|
|
card=persona.P1,
|
|
state=_initial_state(),
|
|
learner_text="Current text has no identifiers.",
|
|
recall_summary=f"Recall mentioned {RAW_PHONE}.",
|
|
pinned_facts=[f"Pinned email {RAW_EMAIL}."],
|
|
recent_turns=[
|
|
{"speaker": "counselor", "text": f"Previous raw RRN {RAW_RRN}."},
|
|
],
|
|
)
|
|
|
|
blob = _message_blob(ctx.messages)
|
|
for raw in RAW_VALUES:
|
|
self.assertNotIn(raw, blob)
|
|
for masked in MASK_VALUES:
|
|
self.assertIn(masked, blob)
|
|
|
|
def test_mask_pii_masks_korean_name_and_institution_context(self) -> None:
|
|
masked = guardrail.mask_pii(
|
|
f"이름: {RAW_KO_NAME}, 소속은 {RAW_KO_ORG} {RAW_KO_DEPT}입니다."
|
|
)
|
|
|
|
self.assertFalse(masked.used_presidio)
|
|
self.assertIn("NAME", masked.entities)
|
|
self.assertIn("ORG", masked.entities)
|
|
for raw in RAW_KO_VALUES:
|
|
self.assertNotIn(raw, masked.text_masked)
|
|
self.assertIn("[NAME]", masked.text_masked)
|
|
self.assertGreaterEqual(masked.text_masked.count("[ORG]"), 2)
|
|
|
|
def test_mask_pii_does_not_mask_common_korean_context_words_as_names(self) -> None:
|
|
masked = guardrail.mask_pii("학교 가는 게 힘들고 엄마랑 친구 이야기를 하면 불안해요.")
|
|
|
|
self.assertEqual(masked.text_masked, "학교 가는 게 힘들고 엄마랑 친구 이야기를 하면 불안해요.")
|
|
self.assertNotIn("NAME", masked.entities)
|
|
self.assertNotIn("ORG", masked.entities)
|
|
|
|
def test_prepare_turn_masks_korean_pii_from_engine_messages(self) -> None:
|
|
ctx = orchestrator.prepare_turn(
|
|
session_id="masking-session",
|
|
case_id="masking-case",
|
|
card=persona.P1,
|
|
state=_initial_state(),
|
|
learner_text=RAW_KO_TEXT,
|
|
recall_summary=f"지난 회기 요약에 {RAW_KO_NAME}과 {RAW_KO_ORG}가 남아 있었다.",
|
|
pinned_facts=[f"소속 {RAW_KO_DEPT}"],
|
|
recent_turns=[
|
|
{"speaker": "counselor", "text": f"{RAW_KO_NAME} 씨가 상담실에 왔다."},
|
|
],
|
|
)
|
|
|
|
blob = _message_blob(ctx.messages)
|
|
_assert_no_raw_ko_pii(self, blob)
|
|
_assert_masked_ko_pii_present(self, blob)
|
|
for raw in RAW_KO_VALUES:
|
|
self.assertIn(raw, ctx.learner_text_raw)
|
|
self.assertNotIn(raw, ctx.learner_text_masked)
|
|
|
|
def test_prepare_turn_threads_theory_mode_into_engine_messages(self) -> None:
|
|
ctx = orchestrator.prepare_turn(
|
|
session_id="theory-session",
|
|
case_id="theory-case",
|
|
card=persona.P2,
|
|
state=_initial_state(),
|
|
learner_text="그냥 아무것도 하기 싫어요.",
|
|
theory_mode="cbt",
|
|
)
|
|
|
|
blob = _message_blob(ctx.messages)
|
|
self.assertIn("[L3-T 이론모드: CBT]", blob)
|
|
self.assertIn("자동적 사고", blob)
|
|
self.assertIn("행동활성화", blob)
|
|
|
|
async def test_run_turn_generate_sends_only_masked_engine_payload(self) -> None:
|
|
ctx = _prepare_context()
|
|
engine = CaptureGenerateEngine()
|
|
audit_payloads: list[dict[str, Any]] = []
|
|
|
|
async def audit_hook(payload: dict[str, Any]) -> None:
|
|
audit_payloads.append(payload)
|
|
|
|
await orchestrator.run_turn_generate(
|
|
ctx,
|
|
engine, # type: ignore[arg-type]
|
|
audit_hook=audit_hook,
|
|
)
|
|
|
|
self.assertIsNotNone(engine.request)
|
|
self.assertIsNotNone(engine.payload)
|
|
_assert_no_raw_pii(self, engine.request.messages)
|
|
_assert_no_raw_pii(self, engine.payload)
|
|
_assert_masked_pii_present(self, engine.request.messages)
|
|
_assert_masked_pii_present(self, engine.payload)
|
|
self.assertEqual(ctx.learner_text_raw, RAW_TEXT)
|
|
self.assertEqual(len(audit_payloads), 1)
|
|
self.assertEqual(audit_payloads[0]["provider"], "fake-provider")
|
|
self.assertEqual(audit_payloads[0]["model"], "fake-model")
|
|
_assert_no_raw_pii(self, audit_payloads)
|
|
for key in ("messages", "prompt", "text"):
|
|
self.assertNotIn(key, audit_payloads[0])
|
|
|
|
async def test_run_turn_stream_sends_only_masked_engine_payload(self) -> None:
|
|
ctx = _prepare_context()
|
|
engine = CaptureStreamEngine()
|
|
audit_payloads: list[dict[str, Any]] = []
|
|
|
|
async def audit_hook(payload: dict[str, Any]) -> None:
|
|
audit_payloads.append(payload)
|
|
|
|
events = [
|
|
event
|
|
async for event in orchestrator.run_turn_stream(
|
|
ctx,
|
|
engine, # type: ignore[arg-type]
|
|
audit_hook=audit_hook,
|
|
)
|
|
]
|
|
|
|
self.assertEqual([event.event for event in events], ["token", "done"])
|
|
self.assertIsNotNone(engine.request)
|
|
self.assertIsNotNone(engine.payload)
|
|
_assert_no_raw_pii(self, engine.request.messages)
|
|
_assert_no_raw_pii(self, engine.payload)
|
|
_assert_masked_pii_present(self, engine.request.messages)
|
|
_assert_masked_pii_present(self, engine.payload)
|
|
self.assertEqual(ctx.learner_text_raw, RAW_TEXT)
|
|
self.assertEqual(len(audit_payloads), 1)
|
|
self.assertEqual(audit_payloads[0]["tokens_in"], 5)
|
|
self.assertEqual(audit_payloads[0]["tokens_out"], 6)
|
|
_assert_no_raw_pii(self, audit_payloads)
|
|
for key in ("messages", "prompt", "text"):
|
|
self.assertNotIn(key, audit_payloads[0])
|
|
|
|
async def test_run_turn_generate_sends_only_masked_korean_pii(self) -> None:
|
|
ctx = orchestrator.prepare_turn(
|
|
session_id="masking-session",
|
|
case_id="masking-case",
|
|
card=persona.P1,
|
|
state=_initial_state(),
|
|
learner_text=RAW_KO_TEXT,
|
|
)
|
|
engine = CaptureGenerateEngine()
|
|
audit_payloads: list[dict[str, Any]] = []
|
|
|
|
async def audit_hook(payload: dict[str, Any]) -> None:
|
|
audit_payloads.append(payload)
|
|
|
|
await orchestrator.run_turn_generate(
|
|
ctx,
|
|
engine, # type: ignore[arg-type]
|
|
audit_hook=audit_hook,
|
|
)
|
|
|
|
self.assertIsNotNone(engine.request)
|
|
self.assertIsNotNone(engine.payload)
|
|
_assert_no_raw_ko_pii(self, engine.request.messages)
|
|
_assert_no_raw_ko_pii(self, engine.payload)
|
|
_assert_masked_ko_pii_present(self, engine.request.messages)
|
|
_assert_masked_ko_pii_present(self, engine.payload)
|
|
_assert_no_raw_ko_pii(self, audit_payloads)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|