Refactor session runtime persistence
This commit is contained in:
@@ -80,6 +80,7 @@ class TestStoredAgentSessionRoundTrip(unittest.TestCase):
|
||||
'user_context': {'lang': 'en'},
|
||||
'system_context': {'os': 'linux'},
|
||||
'messages': ({'role': 'user', 'content': 'hi'},),
|
||||
'display_messages': ({'role': 'user', 'content': 'visible hi'},),
|
||||
'turns': 3,
|
||||
'tool_calls': 7,
|
||||
'usage': {'input_tokens': 500, 'output_tokens': 300},
|
||||
@@ -106,6 +107,7 @@ class TestStoredAgentSessionRoundTrip(unittest.TestCase):
|
||||
self.assertEqual(loaded.user_context, session.user_context)
|
||||
self.assertEqual(loaded.system_context, session.system_context)
|
||||
self.assertEqual(loaded.messages, session.messages)
|
||||
self.assertEqual(loaded.display_messages, session.display_messages)
|
||||
self.assertEqual(loaded.turns, session.turns)
|
||||
self.assertEqual(loaded.tool_calls, session.tool_calls)
|
||||
self.assertEqual(loaded.usage, session.usage)
|
||||
@@ -157,6 +159,30 @@ class TestStoredAgentSessionRoundTrip(unittest.TestCase):
|
||||
self.assertEqual(len(loaded.messages), 2)
|
||||
self.assertEqual(loaded.messages[0]['role'], 'user')
|
||||
self.assertEqual(loaded.messages[1]['role'], 'assistant')
|
||||
self.assertEqual(loaded.display_messages, loaded.messages)
|
||||
|
||||
def test_load_agent_session_falls_back_to_messages_for_display(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
directory = Path(td)
|
||||
path = directory / 'legacy.json'
|
||||
data = {
|
||||
'session_id': 'legacy',
|
||||
'model_config': {},
|
||||
'runtime_config': {'cwd': '/'},
|
||||
'system_prompt_parts': [],
|
||||
'user_context': {},
|
||||
'system_context': {},
|
||||
'messages': [{'role': 'user', 'content': 'legacy hi'}],
|
||||
'turns': 0,
|
||||
'tool_calls': 0,
|
||||
'usage': {},
|
||||
'total_cost_usd': 0.0,
|
||||
'file_history': [],
|
||||
}
|
||||
path.write_text(json.dumps(data))
|
||||
loaded = load_agent_session('legacy', directory=directory)
|
||||
|
||||
self.assertEqual(loaded.display_messages, loaded.messages)
|
||||
|
||||
def test_load_defaults_for_missing_optional_fields(self) -> None:
|
||||
"""Missing optional fields get sensible defaults."""
|
||||
|
||||
Reference in New Issue
Block a user