test_scheduler_broadcast.py 2.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869
  1. import unittest
  2. from unittest.mock import MagicMock, patch, AsyncMock
  3. from app.services.forum_scheduler import ForumScheduler
  4. from app.agent.agent import ParticipantAgent
  5. class TestForumScheduler(unittest.TestCase):
  6. def setUp(self):
  7. self.scheduler = ForumScheduler()
  8. # Mock manager
  9. self.patcher = patch('app.services.forum_scheduler.manager')
  10. self.mock_manager = self.patcher.start()
  11. self.mock_manager.broadcast = AsyncMock()
  12. def tearDown(self):
  13. self.patcher.stop()
  14. def test_agent_speak_broadcasts_chunks(self):
  15. # Setup
  16. mock_db = MagicMock()
  17. forum_id = 1
  18. agent = ParticipantAgent("Test Agent", {"system_prompt": "test"}, 1, "test")
  19. thought = {"action": "speak"}
  20. context = "test context"
  21. # Mock speak to return a generator of chunks
  22. def mock_speak(*args):
  23. return iter("Hello World")
  24. agent.speak = mock_speak
  25. # Mock participants query
  26. mock_p = MagicMock()
  27. mock_p.persona.name = "Test Agent"
  28. mock_p.persona_id = 123
  29. with patch('app.services.forum_scheduler.get_forum_participants', return_value=[mock_p]), \
  30. patch('app.services.forum_scheduler.create_message') as mock_create_msg, \
  31. patch.object(self.scheduler, '_is_forum_running', return_value=True):
  32. mock_create_msg.return_value.id = 1
  33. # Run
  34. import asyncio
  35. asyncio.run(self.scheduler._agent_speak(forum_id, agent, thought, context))
  36. # Verify broadcasts
  37. # We expect len("Hello World") calls to broadcast_chunk
  38. # And 1 call to broadcast_message
  39. # Check broadcast_chunk calls (via manager.broadcast)
  40. # manager.broadcast is called for chunks AND final message
  41. calls = self.mock_manager.broadcast.call_args_list
  42. # Filter for chunks
  43. chunk_calls = [c for c in calls if c[0][1]['type'] == 'message_chunk']
  44. self.assertEqual(len(chunk_calls), len("Hello World"))
  45. # Check content of first chunk
  46. self.assertEqual(chunk_calls[0][0][1]['data']['content'], 'H')
  47. self.assertEqual(chunk_calls[0][0][1]['data']['speaker_name'], "Test Agent")
  48. self.assertEqual(chunk_calls[0][0][1]['data']['persona_id'], 123)
  49. # Filter for final message
  50. final_calls = [c for c in calls if c[0][1]['type'] == 'new_message']
  51. self.assertEqual(len(final_calls), 1)
  52. self.assertEqual(final_calls[0][0][1]['data']['content'], "Hello World")
  53. if __name__ == '__main__':
  54. unittest.main()