from __future__ import annotations import asyncio import unittest from types import SimpleNamespace from unittest.mock import AsyncMock, patch from services.conversation_history import ConversationRecorder class ConversationRecorderTest(unittest.IsolatedAsyncioTestCase): async def test_finish_waits_for_database_cleanup_when_cancelled(self): recorder = ConversationRecorder("conv_test") started = asyncio.Event() release = asyncio.Event() completed = asyncio.Event() async def finish_database_write(*, status: str): self.assertEqual(status, "failed") started.set() await release.wait() completed.set() recorder._finish = finish_database_write task = asyncio.create_task(recorder.finish(status="failed")) await started.wait() task.cancel() await asyncio.sleep(0) self.assertFalse(task.done()) release.set() with self.assertRaises(asyncio.CancelledError): await task self.assertTrue(completed.is_set()) async def test_fixed_reply_transcript_keeps_workflow_metadata(self): recorder = ConversationRecorder("conv_test") recorder._append = AsyncMock() await recorder.record_transport_message( { "type": "transcript", "role": "assistant", "content": "请稍等,我正在处理。", "timestamp": "2026-07-14T10:00:00+08:00", "source": "workflow-fixed-reply", "nodeId": "agent_service", } ) recorder._append.assert_awaited_once_with( "assistant", "请稍等,我正在处理。", "2026-07-14T10:00:00+08:00", { "source": "workflow-fixed-reply", "node_id": "agent_service", }, ) async def test_workflow_trace_is_persisted_without_counting_as_message(self): conversation = SimpleNamespace( extra={"workflow": {"revision": "sha256:test"}}, message_count=3, ) class FakeSession: committed = False async def __aenter__(self): return self async def __aexit__(self, *_args): return None async def get(self, _model, _session_id): return conversation async def commit(self): self.committed = True session = FakeSession() recorder = ConversationRecorder("conv_test") event = { "type": "workflow-event", "eventId": "wfe_test", "event": "node_entered", "timestamp": "2026-08-01T10:00:00+08:00", "workflowRevision": "sha256:test", "transitionId": 0, "nodeId": "start", } with patch( "services.conversation_history.SessionLocal", return_value=session, ): await recorder.record_transport_message(event) await recorder.record_transport_message(event) self.assertTrue(session.committed) self.assertEqual(conversation.message_count, 3) self.assertEqual(len(conversation.extra["workflowTrace"]), 1) saved = conversation.extra["workflowTrace"][0] self.assertEqual(saved["sessionId"], "conv_test") self.assertEqual(saved["sequence"], 1) async def test_image_is_persisted_as_message_with_artifact(self): conversation = SimpleNamespace(message_count=0) class FakeSession: def __init__(self): self.added = [] self.committed = False async def __aenter__(self): return self async def __aexit__(self, *_args): return None def add(self, value): self.added.append(value) async def get(self, _model, _session_id): return conversation async def commit(self): self.committed = True session = FakeSession() recorder = ConversationRecorder("conv_test") with ( patch("services.conversation_history.SessionLocal", return_value=session), patch("services.conversation_history.put_object") as put_object, ): await recorder._record_image( b"jpeg-data", input_id="input_photo", timestamp="2026-08-05T10:00:00+08:00", mime_type="image/jpeg", ) self.assertTrue(session.committed) self.assertEqual(len(session.added), 2) message, artifact = session.added self.assertEqual(message.content_type, "image") self.assertEqual(message.role, "user") self.assertEqual(message.extra["input_id"], "input_photo") self.assertEqual(artifact.message_id, message.id) self.assertEqual(artifact.kind, "image") self.assertEqual(artifact.size_bytes, len(b"jpeg-data")) self.assertEqual(conversation.message_count, 1) put_object.assert_called_once() if __name__ == "__main__": unittest.main()