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.

run_worker_pub_sub.py 2.9 kB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394
  1. import asyncio
  2. import logging
  3. from dataclasses import dataclass
  4. from typing import Any, NoReturn
  5. from agnext.application import WorkerAgentRuntime
  6. from agnext.components import DefaultSubscription, DefaultTopicId, RoutedAgent, message_handler
  7. from agnext.core import MESSAGE_TYPE_REGISTRY, MessageContext
  8. @dataclass
  9. class AskToGreet:
  10. content: str
  11. @dataclass
  12. class Greeting:
  13. content: str
  14. @dataclass
  15. class ReturnedGreeting:
  16. content: str
  17. @dataclass
  18. class Feedback:
  19. content: str
  20. @dataclass
  21. class ReturnedFeedback:
  22. content: str
  23. class ReceiveAgent(RoutedAgent):
  24. def __init__(self) -> None:
  25. super().__init__("Receive Agent")
  26. @message_handler
  27. async def on_greet(self, message: Greeting, ctx: MessageContext) -> None:
  28. await self.publish_message(ReturnedGreeting(f"Returned greeting: {message.content}"), topic_id=DefaultTopicId())
  29. @message_handler
  30. async def on_feedback(self, message: Feedback, ctx: MessageContext) -> None:
  31. await self.publish_message(ReturnedFeedback(f"Returned feedback: {message.content}"), topic_id=DefaultTopicId())
  32. async def on_unhandled_message(self, message: Any, ctx: MessageContext) -> NoReturn: # type: ignore
  33. print(f"Unhandled message: {message}")
  34. class GreeterAgent(RoutedAgent):
  35. def __init__(self) -> None:
  36. super().__init__("Greeter Agent")
  37. @message_handler
  38. async def on_ask(self, message: AskToGreet, ctx: MessageContext) -> None:
  39. await self.publish_message(Greeting(f"Hello, {message.content}!"), topic_id=DefaultTopicId())
  40. @message_handler
  41. async def on_returned_greet(self, message: ReturnedGreeting, ctx: MessageContext) -> None:
  42. await self.publish_message(Feedback(f"Feedback: {message.content}"), topic_id=DefaultTopicId())
  43. async def on_unhandled_message(self, message: Any, ctx: MessageContext) -> NoReturn: # type: ignore
  44. print(f"Unhandled message: {message}")
  45. async def main() -> None:
  46. runtime = WorkerAgentRuntime()
  47. MESSAGE_TYPE_REGISTRY.add_type(Greeting)
  48. MESSAGE_TYPE_REGISTRY.add_type(AskToGreet)
  49. MESSAGE_TYPE_REGISTRY.add_type(Feedback)
  50. MESSAGE_TYPE_REGISTRY.add_type(ReturnedGreeting)
  51. MESSAGE_TYPE_REGISTRY.add_type(ReturnedFeedback)
  52. await runtime.start(host_connection_string="localhost:50051")
  53. await runtime.register("receiver", ReceiveAgent, lambda: [DefaultSubscription()])
  54. await runtime.register("greeter", GreeterAgent, lambda: [DefaultSubscription()])
  55. await runtime.publish_message(AskToGreet("Hello World!"), topic_id=DefaultTopicId())
  56. # Just to keep the runtime running
  57. try:
  58. await asyncio.sleep(1000000)
  59. except KeyboardInterrupt:
  60. pass
  61. await runtime.stop()
  62. if __name__ == "__main__":
  63. logging.basicConfig(level=logging.DEBUG)
  64. logger = logging.getLogger("agnext")
  65. logger.setLevel(logging.DEBUG)
  66. asyncio.run(main())