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

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