event_agent.py 7.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240
  1. import asyncio
  2. import json
  3. from collections.abc import AsyncIterator, Iterable, Sequence
  4. from dataclasses import dataclass
  5. from typing import Any, Protocol
  6. from agent_lab.application.contracts import AgentParams
  7. from agent_lab.application.events import (
  8. EventDefinition,
  9. EventExecutionContext,
  10. EventKernel,
  11. EventRequest,
  12. ResolvedEventArguments,
  13. )
  14. from agent_lab.application.tools import (
  15. ToolExecutionContext,
  16. ToolRegistry,
  17. build_default_tool_registry,
  18. )
  19. from agent_lab.domain.events import ToolCallEvent
  20. from agent_lab.domain.messages import ChatMessage, StreamItem
  21. class EventAgentChatClient(Protocol):
  22. async def stream_chat(
  23. self,
  24. messages: list[ChatMessage],
  25. tools: list[dict[str, Any]],
  26. params: AgentParams,
  27. tool_choice: dict[str, Any] | None = None,
  28. ) -> AsyncIterator[StreamItem]:
  29. ...
  30. @dataclass(frozen=True)
  31. class EventAgentRequest:
  32. events: list[ToolCallEvent]
  33. history: list[ChatMessage]
  34. system_prompt: str = ""
  35. extra_body: dict[str, Any] | None = None
  36. session_id: str | None = None
  37. turn_index: int | None = None
  38. round_index: int | None = None
  39. turn_started_at: float | None = None
  40. class EventAgent:
  41. def __init__(
  42. self,
  43. enabled_tools: Iterable[str],
  44. registry: ToolRegistry | None = None,
  45. chat_client: EventAgentChatClient | None = None,
  46. params: AgentParams | None = None,
  47. ) -> None:
  48. self.enabled_tools = set(enabled_tools)
  49. self.registry = registry or build_default_tool_registry()
  50. self.chat_client = chat_client
  51. self.params = params or AgentParams()
  52. self.kernel = EventKernel(
  53. self.registry.event_registry,
  54. argument_fallback=(
  55. self._resolve_arguments_with_llm if chat_client is not None else None
  56. ),
  57. )
  58. self._raw_chunks_by_event: dict[str, list[dict[str, Any]]] = {}
  59. async def handle(
  60. self,
  61. event: ToolCallEvent,
  62. history: Sequence[ChatMessage] = (),
  63. system_prompt: str = "",
  64. extra_body: dict[str, Any] | None = None,
  65. ) -> ChatMessage:
  66. self._raw_chunks_by_event[event.id] = []
  67. result = await self.kernel.execute(
  68. self.registry.event_request(event),
  69. enabled_names=self.enabled_tools,
  70. context=ToolExecutionContext(
  71. history=history,
  72. system_prompt=system_prompt,
  73. extra_body=(
  74. self.params.extra_body if extra_body is None else extra_body
  75. ),
  76. ),
  77. )
  78. payload = self.registry.tool_payload(result)
  79. return self._tool_reply(event, payload)
  80. async def handle_many(
  81. self,
  82. events: Sequence[ToolCallEvent],
  83. history: Sequence[ChatMessage],
  84. system_prompt: str = "",
  85. extra_body: dict[str, Any] | None = None,
  86. ) -> list[ChatMessage]:
  87. for event in events:
  88. self._raw_chunks_by_event[event.id] = []
  89. return await asyncio.gather(
  90. *[
  91. self.handle(
  92. event,
  93. history=history,
  94. system_prompt=system_prompt,
  95. extra_body=extra_body,
  96. )
  97. for event in events
  98. ]
  99. )
  100. def raw_model_chunks(
  101. self,
  102. events: Sequence[ToolCallEvent],
  103. ) -> list[dict[str, Any]]:
  104. return [
  105. {
  106. "event_id": event.id,
  107. "event_name": event.name,
  108. "chunks": self._raw_chunks_by_event.get(event.id, []),
  109. }
  110. for event in events
  111. ]
  112. async def _resolve_arguments_with_llm(
  113. self,
  114. definition: EventDefinition,
  115. request: EventRequest,
  116. context: EventExecutionContext,
  117. ) -> ResolvedEventArguments | None:
  118. assert self.chat_client is not None
  119. tool = self.registry.tool_schema(request.name)
  120. if tool is None:
  121. return None
  122. event = self.registry._tool_event(request)
  123. messages = self._build_argument_messages(
  124. event,
  125. context.history,
  126. context.system_prompt,
  127. )
  128. params = self.params.model_copy(
  129. update={
  130. "extra_body": (
  131. context.extra_body
  132. if context.extra_body is not None
  133. else self.params.extra_body
  134. ),
  135. }
  136. )
  137. raw_chunks: list[dict[str, Any]] = []
  138. async for item in self.chat_client.stream_chat(
  139. messages=messages,
  140. tools=[tool],
  141. params=params,
  142. tool_choice={
  143. "type": "function",
  144. "function": {"name": event.name},
  145. },
  146. ):
  147. if item.kind == "raw_chunk" and item.raw_chunk is not None:
  148. raw_chunks.append(item.raw_chunk)
  149. continue
  150. if item.kind == "provider_tool_call" and item.event is not None:
  151. self._raw_chunks_by_event[event.id] = raw_chunks
  152. return ResolvedEventArguments(
  153. event_name=item.event.name,
  154. arguments=dict(item.event.arguments),
  155. raw_arguments=item.event.raw_arguments,
  156. )
  157. self._raw_chunks_by_event[event.id] = raw_chunks
  158. return None
  159. def _build_argument_messages(
  160. self,
  161. event: ToolCallEvent,
  162. history: Sequence[ChatMessage],
  163. system_prompt: str,
  164. ) -> list[ChatMessage]:
  165. messages = [
  166. ChatMessage(
  167. role="system",
  168. content=(
  169. "Generate parameters for the requested EventAgent tool. "
  170. "Call exactly the provided tool with valid JSON arguments."
  171. ),
  172. )
  173. ]
  174. if system_prompt.strip():
  175. messages.append(ChatMessage(role="system", content=system_prompt))
  176. messages.extend(self._project_visible_history(history))
  177. messages.append(
  178. ChatMessage(
  179. role="user",
  180. content=f"Resolve arguments for event `{event.name}`.",
  181. )
  182. )
  183. return messages
  184. def _project_visible_history(
  185. self,
  186. history: Sequence[ChatMessage],
  187. ) -> list[ChatMessage]:
  188. projected: list[ChatMessage] = []
  189. for message in history:
  190. if message.role == "tool":
  191. continue
  192. if message.role == "assistant":
  193. if not message.content:
  194. continue
  195. projected.append(
  196. ChatMessage(
  197. role="assistant",
  198. content=message.content,
  199. name=message.name,
  200. )
  201. )
  202. continue
  203. projected.append(message)
  204. return projected
  205. def summarize_replies(self, replies: Sequence[ChatMessage]) -> ChatMessage | None:
  206. if not replies:
  207. return None
  208. return ChatMessage(
  209. role="user",
  210. content=(
  211. "EventAgent results:\n"
  212. + "\n".join(reply.content for reply in replies)
  213. ),
  214. name="event_agent",
  215. )
  216. def _tool_reply(self, event: ToolCallEvent, payload: dict[str, Any]) -> ChatMessage:
  217. return ChatMessage(
  218. role="tool",
  219. content=json.dumps(payload, ensure_ascii=False),
  220. name=event.name,
  221. tool_call_id=event.id,
  222. )