68 lines
2.1 KiB
Python
68 lines
2.1 KiB
Python
import unittest
|
|
from contextlib import contextmanager
|
|
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)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|