You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

test_intervention.py 4.8 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127
  1. import pytest
  2. from autogen_core import AgentId, DefaultInterventionHandler, DropMessage, SingleThreadedAgentRuntime
  3. from autogen_core.exceptions import MessageDroppedException
  4. from autogen_test_utils import LoopbackAgent, MessageType
  5. @pytest.mark.asyncio
  6. async def test_intervention_count_messages() -> None:
  7. class DebugInterventionHandler(DefaultInterventionHandler):
  8. def __init__(self) -> None:
  9. self.num_messages = 0
  10. async def on_send(self, message: MessageType, *, sender: AgentId | None, recipient: AgentId) -> MessageType:
  11. self.num_messages += 1
  12. return message
  13. handler = DebugInterventionHandler()
  14. runtime = SingleThreadedAgentRuntime(intervention_handlers=[handler])
  15. await LoopbackAgent.register(runtime, "name", LoopbackAgent)
  16. loopback = AgentId("name", key="default")
  17. runtime.start()
  18. _response = await runtime.send_message(MessageType(), recipient=loopback)
  19. await runtime.stop()
  20. assert handler.num_messages == 1
  21. loopback_agent = await runtime.try_get_underlying_agent_instance(loopback, type=LoopbackAgent)
  22. assert loopback_agent.num_calls == 1
  23. @pytest.mark.asyncio
  24. async def test_intervention_drop_send() -> None:
  25. class DropSendInterventionHandler(DefaultInterventionHandler):
  26. async def on_send(
  27. self, message: MessageType, *, sender: AgentId | None, recipient: AgentId
  28. ) -> MessageType | type[DropMessage]:
  29. return DropMessage
  30. handler = DropSendInterventionHandler()
  31. runtime = SingleThreadedAgentRuntime(intervention_handlers=[handler])
  32. await LoopbackAgent.register(runtime, "name", LoopbackAgent)
  33. loopback = AgentId("name", key="default")
  34. runtime.start()
  35. with pytest.raises(MessageDroppedException):
  36. _response = await runtime.send_message(MessageType(), recipient=loopback)
  37. await runtime.stop()
  38. loopback_agent = await runtime.try_get_underlying_agent_instance(loopback, type=LoopbackAgent)
  39. assert loopback_agent.num_calls == 0
  40. @pytest.mark.asyncio
  41. async def test_intervention_drop_response() -> None:
  42. class DropResponseInterventionHandler(DefaultInterventionHandler):
  43. async def on_response(
  44. self, message: MessageType, *, sender: AgentId, recipient: AgentId | None
  45. ) -> MessageType | type[DropMessage]:
  46. return DropMessage
  47. handler = DropResponseInterventionHandler()
  48. runtime = SingleThreadedAgentRuntime(intervention_handlers=[handler])
  49. await LoopbackAgent.register(runtime, "name", LoopbackAgent)
  50. loopback = AgentId("name", key="default")
  51. runtime.start()
  52. with pytest.raises(MessageDroppedException):
  53. _response = await runtime.send_message(MessageType(), recipient=loopback)
  54. await runtime.stop()
  55. @pytest.mark.asyncio
  56. async def test_intervention_raise_exception_on_send() -> None:
  57. class InterventionException(Exception):
  58. pass
  59. class ExceptionInterventionHandler(DefaultInterventionHandler): # type: ignore
  60. async def on_send(
  61. self, message: MessageType, *, sender: AgentId | None, recipient: AgentId
  62. ) -> MessageType | type[DropMessage]: # type: ignore
  63. raise InterventionException
  64. handler = ExceptionInterventionHandler()
  65. runtime = SingleThreadedAgentRuntime(intervention_handlers=[handler])
  66. await LoopbackAgent.register(runtime, "name", LoopbackAgent)
  67. loopback = AgentId("name", key="default")
  68. runtime.start()
  69. with pytest.raises(InterventionException):
  70. _response = await runtime.send_message(MessageType(), recipient=loopback)
  71. await runtime.stop()
  72. long_running_agent = await runtime.try_get_underlying_agent_instance(loopback, type=LoopbackAgent)
  73. assert long_running_agent.num_calls == 0
  74. @pytest.mark.asyncio
  75. async def test_intervention_raise_exception_on_respond() -> None:
  76. class InterventionException(Exception):
  77. pass
  78. class ExceptionInterventionHandler(DefaultInterventionHandler): # type: ignore
  79. async def on_response(
  80. self, message: MessageType, *, sender: AgentId, recipient: AgentId | None
  81. ) -> MessageType | type[DropMessage]: # type: ignore
  82. raise InterventionException
  83. handler = ExceptionInterventionHandler()
  84. runtime = SingleThreadedAgentRuntime(intervention_handlers=[handler])
  85. await LoopbackAgent.register(runtime, "name", LoopbackAgent)
  86. loopback = AgentId("name", key="default")
  87. runtime.start()
  88. with pytest.raises(InterventionException):
  89. _response = await runtime.send_message(MessageType(), recipient=loopback)
  90. await runtime.stop()
  91. long_running_agent = await runtime.try_get_underlying_agent_instance(loopback, type=LoopbackAgent)
  92. assert long_running_agent.num_calls == 1