141 lines
4.5 KiB
Python
141 lines
4.5 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),
|
|
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()
|