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.

_single_threaded_agent_runtime.py 11 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294
  1. import asyncio
  2. import logging
  3. from asyncio import Future
  4. from collections.abc import Sequence
  5. from dataclasses import dataclass
  6. from typing import Any, Awaitable, Dict, List, Mapping, Set
  7. from ..core import Agent, AgentRuntime, CancellationToken
  8. from ..core.exceptions import MessageDroppedException
  9. from ..core.intervention import DropMessage, InterventionHandler
  10. logger = logging.getLogger("agnext")
  11. event_logger = logging.getLogger("agnext.events")
  12. @dataclass(kw_only=True)
  13. class PublishMessageEnvelope:
  14. """A message envelope for publishing messages to all agents that can handle
  15. the message of the type T."""
  16. message: Any
  17. cancellation_token: CancellationToken
  18. sender: Agent | None
  19. @dataclass(kw_only=True)
  20. class SendMessageEnvelope:
  21. """A message envelope for sending a message to a specific agent that can handle
  22. the message of the type T."""
  23. message: Any
  24. sender: Agent | None
  25. recipient: Agent
  26. future: Future[Any]
  27. cancellation_token: CancellationToken
  28. @dataclass(kw_only=True)
  29. class ResponseMessageEnvelope:
  30. """A message envelope for sending a response to a message."""
  31. message: Any
  32. future: Future[Any]
  33. sender: Agent
  34. recipient: Agent | None
  35. class SingleThreadedAgentRuntime(AgentRuntime):
  36. def __init__(self, *, before_send: InterventionHandler | None = None) -> None:
  37. self._message_queue: List[PublishMessageEnvelope | SendMessageEnvelope | ResponseMessageEnvelope] = []
  38. self._per_type_subscribers: Dict[type, List[Agent]] = {}
  39. self._agents: Set[Agent] = set()
  40. self._before_send = before_send
  41. def add_agent(self, agent: Agent) -> None:
  42. agent_names = {agent.name for agent in self._agents}
  43. if agent.name in agent_names:
  44. raise ValueError(f"Agent with name {agent.name} already exists. Agent names must be unique.")
  45. for message_type in agent.subscriptions:
  46. if message_type not in self._per_type_subscribers:
  47. self._per_type_subscribers[message_type] = []
  48. self._per_type_subscribers[message_type].append(agent)
  49. self._agents.add(agent)
  50. @property
  51. def agents(self) -> Sequence[Agent]:
  52. return list(self._agents)
  53. @property
  54. def unprocessed_messages(
  55. self,
  56. ) -> Sequence[PublishMessageEnvelope | SendMessageEnvelope | ResponseMessageEnvelope]:
  57. return self._message_queue
  58. # Returns the response of the message
  59. def send_message(
  60. self,
  61. message: Any,
  62. recipient: Agent,
  63. *,
  64. sender: Agent | None = None,
  65. cancellation_token: CancellationToken | None = None,
  66. ) -> Future[Any | None]:
  67. if cancellation_token is None:
  68. cancellation_token = CancellationToken()
  69. logger.info(f"Sending message of type {type(message).__name__} to {recipient.name}: {message.__dict__}")
  70. # event_logger.info(
  71. # MessageEvent(
  72. # payload=message,
  73. # sender=sender,
  74. # receiver=recipient,
  75. # kind=MessageKind.DIRECT,
  76. # delivery_stage=DeliveryStage.SEND,
  77. # )
  78. # )
  79. future = asyncio.get_event_loop().create_future()
  80. if recipient not in self._agents:
  81. future.set_exception(Exception("Recipient not found"))
  82. self._message_queue.append(
  83. SendMessageEnvelope(
  84. message=message,
  85. recipient=recipient,
  86. future=future,
  87. cancellation_token=cancellation_token,
  88. sender=sender,
  89. )
  90. )
  91. return future
  92. def publish_message(
  93. self,
  94. message: Any,
  95. *,
  96. sender: Agent | None = None,
  97. cancellation_token: CancellationToken | None = None,
  98. ) -> Future[None]:
  99. if cancellation_token is None:
  100. cancellation_token = CancellationToken()
  101. logger.info(f"Publishing message of type {type(message).__name__} to all subscribers: {message.__dict__}")
  102. # event_logger.info(
  103. # MessageEvent(
  104. # payload=message,
  105. # sender=sender,
  106. # receiver=None,
  107. # kind=MessageKind.PUBLISH,
  108. # delivery_stage=DeliveryStage.SEND,
  109. # )
  110. # )
  111. self._message_queue.append(
  112. PublishMessageEnvelope(
  113. message=message,
  114. cancellation_token=cancellation_token,
  115. sender=sender,
  116. )
  117. )
  118. future = asyncio.get_event_loop().create_future()
  119. future.set_result(None)
  120. return future
  121. def save_state(self) -> Mapping[str, Any]:
  122. state: Dict[str, Dict[str, Any]] = {}
  123. for agent in self._agents:
  124. state[agent.name] = dict(agent.save_state())
  125. return state
  126. def load_state(self, state: Mapping[str, Any]) -> None:
  127. for agent in self._agents:
  128. agent.load_state(state[agent.name])
  129. async def _process_send(self, message_envelope: SendMessageEnvelope) -> None:
  130. recipient = message_envelope.recipient
  131. assert recipient in self._agents
  132. try:
  133. sender_name = message_envelope.sender.name if message_envelope.sender is not None else "Unknown"
  134. logger.info(
  135. f"Calling message handler for {recipient.name} with message type {type(message_envelope.message).__name__} sent by {sender_name}"
  136. )
  137. # event_logger.info(
  138. # MessageEvent(
  139. # payload=message_envelope.message,
  140. # sender=message_envelope.sender,
  141. # receiver=recipient,
  142. # kind=MessageKind.DIRECT,
  143. # delivery_stage=DeliveryStage.DELIVER,
  144. # )
  145. # )
  146. response = await recipient.on_message(
  147. message_envelope.message,
  148. cancellation_token=message_envelope.cancellation_token,
  149. )
  150. except BaseException as e:
  151. message_envelope.future.set_exception(e)
  152. return
  153. self._message_queue.append(
  154. ResponseMessageEnvelope(
  155. message=response,
  156. future=message_envelope.future,
  157. sender=message_envelope.recipient,
  158. recipient=message_envelope.sender,
  159. )
  160. )
  161. async def _process_publish(self, message_envelope: PublishMessageEnvelope) -> None:
  162. responses: List[Awaitable[Any]] = []
  163. for agent in self._per_type_subscribers.get(type(message_envelope.message), []): # type: ignore
  164. if message_envelope.sender is not None and agent.name == message_envelope.sender.name:
  165. continue
  166. sender_name = message_envelope.sender.name if message_envelope.sender is not None else "Unknown"
  167. logger.info(
  168. f"Calling message handler for {agent.name} with message type {type(message_envelope.message).__name__} published by {sender_name}"
  169. )
  170. # event_logger.info(
  171. # MessageEvent(
  172. # payload=message_envelope.message,
  173. # sender=message_envelope.sender,
  174. # receiver=agent,
  175. # kind=MessageKind.PUBLISH,
  176. # delivery_stage=DeliveryStage.DELIVER,
  177. # )
  178. # )
  179. future = agent.on_message(
  180. message_envelope.message,
  181. cancellation_token=message_envelope.cancellation_token,
  182. )
  183. responses.append(future)
  184. try:
  185. _all_responses = await asyncio.gather(*responses)
  186. except BaseException:
  187. logger.error("Error processing publish message", exc_info=True)
  188. return
  189. # TODO if responses are given for a publish
  190. async def _process_response(self, message_envelope: ResponseMessageEnvelope) -> None:
  191. recipient_name = message_envelope.recipient.name if message_envelope.recipient is not None else "Unknown"
  192. content = (
  193. message_envelope.message.__dict__
  194. if hasattr(message_envelope.message, "__dict__")
  195. else message_envelope.message
  196. )
  197. logger.info(
  198. f"Resolving response with message type {type(message_envelope.message).__name__} for recipient {recipient_name} from {message_envelope.sender.name}: {content}"
  199. )
  200. # event_logger.info(
  201. # MessageEvent(
  202. # payload=message_envelope.message,
  203. # sender=message_envelope.sender,
  204. # receiver=message_envelope.recipient,
  205. # kind=MessageKind.RESPOND,
  206. # delivery_stage=DeliveryStage.DELIVER,
  207. # )
  208. # )
  209. message_envelope.future.set_result(message_envelope.message)
  210. async def process_next(self) -> None:
  211. if len(self._message_queue) == 0:
  212. # Yield control to the event loop to allow other tasks to run
  213. await asyncio.sleep(0)
  214. return
  215. message_envelope = self._message_queue.pop(0)
  216. match message_envelope:
  217. case SendMessageEnvelope(message=message, sender=sender, recipient=recipient, future=future):
  218. if self._before_send is not None:
  219. temp_message = await self._before_send.on_send(message, sender=sender, recipient=recipient)
  220. if temp_message is DropMessage or isinstance(temp_message, DropMessage):
  221. future.set_exception(MessageDroppedException())
  222. return
  223. message_envelope.message = temp_message
  224. asyncio.create_task(self._process_send(message_envelope))
  225. case PublishMessageEnvelope(
  226. message=message,
  227. sender=sender,
  228. ):
  229. if self._before_send is not None:
  230. temp_message = await self._before_send.on_publish(message, sender=sender)
  231. if temp_message is DropMessage or isinstance(temp_message, DropMessage):
  232. # TODO log message dropped
  233. return
  234. message_envelope.message = temp_message
  235. asyncio.create_task(self._process_publish(message_envelope))
  236. case ResponseMessageEnvelope(message=message, sender=sender, recipient=recipient, future=future):
  237. if self._before_send is not None:
  238. temp_message = await self._before_send.on_response(message, sender=sender, recipient=recipient)
  239. if temp_message is DropMessage or isinstance(temp_message, DropMessage):
  240. future.set_exception(MessageDroppedException())
  241. return
  242. message_envelope.message = temp_message
  243. asyncio.create_task(self._process_response(message_envelope))
  244. # Yield control to the message loop to allow other tasks to run
  245. await asyncio.sleep(0)

This is a mirror of AutoGen from GitHub. AutoGen is a framework that enables the development of LLM applications using multiple agents that can converse with each other to solve tasks.