test_stream_robustness.py 3.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980
  1. import unittest
  2. from unittest.mock import MagicMock, patch, AsyncMock
  3. import asyncio
  4. import json
  5. from app.services.forum_scheduler import ForumScheduler
  6. from app.agent.agent import ModeratorAgent
  7. class TestStreamRobustness(unittest.IsolatedAsyncioTestCase):
  8. async def test_moderator_stream_fields(self):
  9. """
  10. Verify that moderator streaming broadcasts include stream_id and moderator_id.
  11. """
  12. scheduler = ForumScheduler()
  13. # Mock DB and objects
  14. mock_db = MagicMock()
  15. mock_forum = MagicMock()
  16. mock_forum.id = 1
  17. mock_forum.moderator_id = 99
  18. mock_forum.summary_history = []
  19. # Mock get_forum to return our mock forum
  20. # We need to patch 'app.services.forum_scheduler.get_forum'
  21. # Mock ModeratorAgent to return a generator
  22. mock_moderator = MagicMock(spec=ModeratorAgent)
  23. mock_moderator.name = "TestHost"
  24. def mock_opening(guests):
  25. yield "Hello"
  26. yield " World"
  27. # Patch dependencies
  28. with patch('app.services.forum_scheduler.get_forum', return_value=mock_forum), \
  29. patch('app.services.forum_scheduler.create_message') as mock_create_msg, \
  30. patch('app.services.forum_scheduler.update_forum'), \
  31. patch('app.services.forum_scheduler.manager') as mock_manager, \
  32. patch.object(scheduler, '_is_forum_running', return_value=True), \
  33. patch('asyncio.to_thread', side_effect=lambda func, *args: func(*args)) as mock_to_thread:
  34. # Make broadcast awaitable
  35. mock_manager.broadcast = AsyncMock()
  36. # Setup moderator mock methods
  37. mock_moderator.opening = mock_opening
  38. # Run _moderator_speak
  39. # We assume asyncio.to_thread executes the function immediately for this test
  40. mock_create_msg.return_value.id = 1
  41. await scheduler._moderator_speak(1, mock_moderator, "opening", guests=[])
  42. # Verify broadcasts
  43. calls = mock_manager.broadcast.call_args_list
  44. chunk_calls = [call for call in calls if call[0][1]['type'] == 'message_chunk']
  45. message_calls = [call for call in calls if call[0][1]['type'] == 'new_message']
  46. speech_logs = [
  47. call for call in calls
  48. if call[0][1]['type'] == 'system_log'
  49. and call[0][1]['data']['level'] == 'speech'
  50. ]
  51. self.assertEqual(len(chunk_calls), 2)
  52. self.assertEqual(len(message_calls), 1)
  53. self.assertGreaterEqual(len(speech_logs), 1)
  54. # Check that stream_id and moderator_id are present in chunks
  55. # First call: Chunk 1
  56. call_args_1 = chunk_calls[0]
  57. payload_1 = call_args_1[0][1]
  58. self.assertEqual(payload_1['type'], 'message_chunk')
  59. self.assertIn('stream_id', payload_1['data'])
  60. self.assertEqual(payload_1['data']['moderator_id'], 99)
  61. call_args_msg = message_calls[0]
  62. payload_msg = call_args_msg[0][1]
  63. self.assertEqual(payload_msg['type'], 'new_message')
  64. self.assertIn('stream_id', payload_msg['data'])
  65. self.assertEqual(payload_msg['data']['moderator_id'], 99)
  66. if __name__ == '__main__':
  67. unittest.main()