85 lines
2.7 KiB
Python
85 lines
2.7 KiB
Python
from __future__ import annotations
|
|
|
|
import unittest
|
|
|
|
from src.agent_session import AgentMessage, sanitize_model_message_sequence
|
|
|
|
|
|
class TestSanitizeModelMessageSequence(unittest.TestCase):
|
|
def test_removes_orphan_tool_message(self) -> None:
|
|
messages = [
|
|
AgentMessage(role='user', content='summary', message_id='summary'),
|
|
AgentMessage(
|
|
role='tool',
|
|
content='orphan result',
|
|
tool_call_id='call_1',
|
|
message_id='tool_1',
|
|
),
|
|
AgentMessage(role='assistant', content='review', message_id='assistant_1'),
|
|
AgentMessage(role='user', content='continue', message_id='user_1'),
|
|
]
|
|
|
|
cleaned = sanitize_model_message_sequence(messages)
|
|
|
|
self.assertEqual(
|
|
[message.message_id for message in cleaned],
|
|
['summary', 'assistant_1', 'user_1'],
|
|
)
|
|
|
|
def test_preserves_complete_assistant_tool_group(self) -> None:
|
|
assistant = AgentMessage(
|
|
role='assistant',
|
|
content='',
|
|
tool_calls=(
|
|
{
|
|
'id': 'call_1',
|
|
'type': 'function',
|
|
'function': {'name': 'read_file', 'arguments': '{}'},
|
|
},
|
|
),
|
|
message_id='assistant_1',
|
|
)
|
|
tool = AgentMessage(
|
|
role='tool',
|
|
content='ok',
|
|
tool_call_id='call_1',
|
|
message_id='tool_1',
|
|
)
|
|
final = AgentMessage(role='assistant', content='done', message_id='assistant_2')
|
|
|
|
cleaned = sanitize_model_message_sequence([assistant, tool, final])
|
|
|
|
self.assertEqual(
|
|
[message.message_id for message in cleaned],
|
|
['assistant_1', 'tool_1', 'assistant_2'],
|
|
)
|
|
|
|
def test_removes_incomplete_assistant_tool_call_group(self) -> None:
|
|
messages = [
|
|
AgentMessage(role='user', content='question', message_id='user_1'),
|
|
AgentMessage(
|
|
role='assistant',
|
|
content='',
|
|
tool_calls=(
|
|
{
|
|
'id': 'call_1',
|
|
'type': 'function',
|
|
'function': {'name': 'read_file', 'arguments': '{}'},
|
|
},
|
|
),
|
|
message_id='assistant_1',
|
|
),
|
|
AgentMessage(role='user', content='next', message_id='user_2'),
|
|
]
|
|
|
|
cleaned = sanitize_model_message_sequence(messages)
|
|
|
|
self.assertEqual(
|
|
[message.message_id for message in cleaned],
|
|
['user_1', 'user_2'],
|
|
)
|
|
|
|
|
|
if __name__ == '__main__':
|
|
unittest.main()
|