zcbot/tests/test_message_payload.py

135 lines
4.3 KiB
Python

import unittest
from contextlib import contextmanager
from types import SimpleNamespace
from unittest.mock import patch
from uuid import uuid4
from core.session import Session
from core.storage.models import Message, sanitize_jsonb_nul
class MessagePayloadSanitizationTests(unittest.TestCase):
def test_recursively_removes_nul_and_preserves_other_text(self) -> None:
payload = {
"role": "tool",
"content": "In 10\tH \x00unnamed\n中文",
"nested\x00key": ["a\x00b", {"value": "c\x00d"}],
}
cleaned = sanitize_jsonb_nul(payload)
self.assertEqual(cleaned["content"], "In 10\tH unnamed\n中文")
self.assertEqual(cleaned["nestedkey"], ["ab", {"value": "cd"}])
self.assertEqual(payload["content"], "In 10\tH \x00unnamed\n中文")
def test_message_model_sanitizes_payload_on_assignment(self) -> None:
message = Message(
task_id=uuid4(),
idx=0,
payload={
"role": "tool",
"content": "before\x00after",
"tool_calls": ({"arguments": "x\x00y"},),
},
)
self.assertEqual(message.payload["content"], "beforeafter")
self.assertEqual(
message.payload["tool_calls"],
({"arguments": "xy"},),
)
def test_session_keeps_sanitized_value_in_memory_and_database_row(self) -> None:
class FakeDbSession:
row: Message
def add(self, row: Message) -> None:
self.row = row
def flush(self) -> None:
self.row.message_id = uuid4()
fake_db = FakeDbSession()
@contextmanager
def fake_session_scope():
yield fake_db
session = Session(task_id=uuid4())
with patch("core.session.session_scope", side_effect=fake_session_scope):
session.append({"role": "tool", "content": "before\x00after"})
expected = {"role": "tool", "content": "beforeafter"}
self.assertEqual(session.messages, [expected])
self.assertEqual(fake_db.row.payload, expected)
def test_session_repairs_assistant_markdown_before_storing_it(self) -> None:
class FakeDbSession:
row: Message
def add(self, row: Message) -> None:
self.row = row
def flush(self) -> None:
self.row.message_id = uuid4()
fake_db = FakeDbSession()
@contextmanager
def fake_session_scope():
yield fake_db
broken = "```markdown\n```mermaid\nA --> B\n```\n```\n正文\n"
session = Session(task_id=uuid4())
with patch("core.session.session_scope", side_effect=fake_session_scope):
session.append({"role": "assistant", "content": broken})
stored = session.messages[0]["content"]
self.assertTrue(stored.startswith("````markdown\n```mermaid"))
self.assertIn("```\n````\n正文", stored)
self.assertEqual(fake_db.row.payload["content"], stored)
def test_session_repairs_historical_assistant_markdown_in_memory_only(self) -> None:
broken = "```markdown\n```mermaid\nA --> B\n```\n```\n正文\n"
class FakeResult:
def __init__(self, value):
self.value = value
def first(self):
return self.value
def scalars(self):
return self
def all(self):
return self.value
def scalar_one(self):
return self.value
class FakeDbSession:
def __init__(self):
self.results = iter([
FakeResult(SimpleNamespace(context_base_idx=0, context_summary="")),
FakeResult([SimpleNamespace(payload={"role": "assistant", "content": broken})]),
FakeResult(1),
])
def execute(self, _query):
return next(self.results)
@contextmanager
def fake_session_scope():
yield FakeDbSession()
with patch("core.session.session_scope", side_effect=fake_session_scope):
session = Session.load(uuid4())
self.assertTrue(session.messages[0]["content"].startswith("````markdown"))
self.assertEqual(session._db_idx, 1)
if __name__ == "__main__":
unittest.main()