Files
ai-video-fullstack/backend/tests/test_conversation_history.py
2026-08-07 10:59:47 +08:00

205 lines
6.6 KiB
Python

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_completed_conversation_queues_enabled_analysis(self):
conversation = SimpleNamespace(
status="active",
ended_at=None,
message_count=0,
analysis_status="none",
analysis_data={},
analysis_error="",
)
class FakeSession:
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):
return None
plan = {
"enabled": True,
"model_resource_id": "model_001",
"fields": [{"name": "customer_name", "type": "string"}],
}
recorder = ConversationRecorder("conv_test", plan)
recorder._sequence = 1
with patch(
"services.conversation_history.SessionLocal",
return_value=FakeSession(),
):
await recorder._finish(status="completed")
self.assertEqual(conversation.analysis_status, "pending")
self.assertEqual(conversation.analysis_data, {"plan": plan})
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 flush(self):
return None
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",
content="帮我看看",
source="uploaded_asset",
)
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.content, "帮我看看")
self.assertEqual(message.extra["input_id"], "input_photo")
self.assertEqual(message.extra["source"], "uploaded_asset")
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()