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), patch("core.session.allocate_message_idx", return_value=0), ): 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), patch("core.session.allocate_message_idx", return_value=0), ): 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()