ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
from aworld.runners.handler.base import DefaultHandler
|
||||
from aworld.utils.common import scan_packages
|
||||
|
||||
scan_packages("aworld.runners.handler", [DefaultHandler])
|
||||
@@ -0,0 +1,481 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
import abc
|
||||
from typing import AsyncGenerator, Tuple
|
||||
|
||||
from aworld.agents.loop_llm_agent import LoopableAgent
|
||||
from aworld.core.agent.base import is_agent, AgentFactory
|
||||
from aworld.core.agent.swarm import GraphBuildType, AgentGraph
|
||||
from aworld.core.common import ActionModel, Observation, TaskItem
|
||||
from aworld.core.event.base import Message, Constants, TopicType, AgentMessage
|
||||
from aworld.core.exceptions import AWorldRuntimeException
|
||||
from aworld.logs.util import logger
|
||||
from aworld.runners import HandlerFactory
|
||||
from aworld.runners.handler.base import DefaultHandler
|
||||
from aworld.runners.handler.tool import DefaultToolHandler
|
||||
from aworld.runners.state_manager import RunNode, RunNodeStatus, RunNodeBusiType
|
||||
from aworld.runners.utils import endless_detect
|
||||
from aworld.output.base import StepOutput
|
||||
|
||||
|
||||
class AgentHandler(DefaultHandler):
|
||||
__metaclass__ = abc.ABCMeta
|
||||
|
||||
def __init__(self, runner: 'TaskEventRunner'):
|
||||
super().__init__(runner)
|
||||
self.runner = runner
|
||||
self.swarm = runner.swarm
|
||||
self.endless_threshold = runner.endless_threshold
|
||||
self.task_id = runner.task.id
|
||||
|
||||
self.agent_calls = []
|
||||
|
||||
@classmethod
|
||||
def name(cls):
|
||||
return "_agents_handler"
|
||||
|
||||
|
||||
@HandlerFactory.register(name=f'__{Constants.AGENT}__')
|
||||
class DefaultAgentHandler(AgentHandler):
|
||||
def is_valid_message(self, message: Message):
|
||||
if message.category != Constants.AGENT:
|
||||
if self.swarm and message.sender in self.swarm.agents and message.sender in AgentFactory:
|
||||
if self.agent_calls:
|
||||
if self.agent_calls[-1] != message.sender:
|
||||
self.agent_calls.append(message.sender)
|
||||
else:
|
||||
self.agent_calls.append(message.sender)
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _do_handle(self, message: Message) -> AsyncGenerator[Message, None]:
|
||||
if not self.is_valid_message(message):
|
||||
return
|
||||
|
||||
headers = {"context": message.context}
|
||||
session_id = message.session_id
|
||||
data = message.payload
|
||||
if not data:
|
||||
# error message, p2p
|
||||
yield Message(
|
||||
category=Constants.OUTPUT,
|
||||
payload=StepOutput.build_failed_output(name=f"{message.caller or self.name()}",
|
||||
step_num=0,
|
||||
data="no data to process.",
|
||||
task_id=self.task_id),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
headers=headers
|
||||
)
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg="no data to process.", data=data, stop=True),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=headers
|
||||
)
|
||||
return
|
||||
|
||||
if isinstance(data, Tuple) and isinstance(data[0], Observation):
|
||||
data = data[0]
|
||||
message.payload = data
|
||||
# data is Observation
|
||||
if isinstance(data, Observation):
|
||||
if not self.swarm:
|
||||
msg = Message(
|
||||
category=Constants.TASK,
|
||||
payload=data.content,
|
||||
sender=data.observer,
|
||||
session_id=session_id,
|
||||
topic=TopicType.FINISHED,
|
||||
headers=headers
|
||||
)
|
||||
logger.info(f"FINISHED|agent handler send finished message: {msg}")
|
||||
yield msg
|
||||
return
|
||||
|
||||
agent = self.swarm.agents.get(message.receiver)
|
||||
# agent + tool completion protocol.
|
||||
if agent and agent.finished and data.info.get('done'):
|
||||
self.swarm.cur_step += 1
|
||||
|
||||
root_agent = self.swarm.communicate_agent
|
||||
if isinstance(root_agent, list):
|
||||
root_agent = root_agent[0]
|
||||
if agent.id() == root_agent.id():
|
||||
msg = Message(
|
||||
category=Constants.TASK,
|
||||
payload=data.content,
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.FINISHED,
|
||||
headers=headers
|
||||
)
|
||||
logger.info(f"FINISHED|agent handler send finished message: {msg}")
|
||||
yield msg
|
||||
else:
|
||||
msg = Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content=data.content),
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
receiver=root_agent.id(),
|
||||
headers=message.headers
|
||||
)
|
||||
logger.info(f"agent handler send agent message: {msg}")
|
||||
yield msg
|
||||
else:
|
||||
if data.info.get('done'):
|
||||
agent_name = self.agent_calls[-1]
|
||||
async for event in self._stop_check(ActionModel(agent_name=agent_name, policy_info=data.content),
|
||||
message):
|
||||
yield event
|
||||
elif not message.receiver:
|
||||
agent_name = message.sender
|
||||
async for event in self._stop_check(ActionModel(agent_name=agent_name, policy_info=data.content),
|
||||
message):
|
||||
yield event
|
||||
else:
|
||||
logger.info(f"agent handler send observation message: {message}")
|
||||
yield message
|
||||
return
|
||||
|
||||
# data is List[ActionModel]
|
||||
for action in data:
|
||||
if not isinstance(action, ActionModel):
|
||||
# error message, p2p
|
||||
yield Message(
|
||||
category=Constants.OUTPUT,
|
||||
payload=StepOutput.build_failed_output(name=f"{message.caller or self.name()}",
|
||||
step_num=0,
|
||||
data="action not a ActionModel.",
|
||||
task_id=self.task_id),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
headers=headers
|
||||
)
|
||||
msg = Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg="action not a ActionModel.", data=data, stop=True),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=headers
|
||||
)
|
||||
logger.info(f"agent handler send task message: {msg}")
|
||||
yield msg
|
||||
return
|
||||
|
||||
tools = []
|
||||
agents = []
|
||||
for action in data:
|
||||
if is_agent(action):
|
||||
agents.append(action)
|
||||
else:
|
||||
tools.append(action)
|
||||
|
||||
if tools:
|
||||
msg = Message(
|
||||
category=Constants.TOOL,
|
||||
payload=tools,
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
receiver=DefaultToolHandler.name(),
|
||||
headers=message.headers
|
||||
)
|
||||
logger.info(f"agent handler send tool message: {msg}")
|
||||
yield msg
|
||||
else:
|
||||
yield Message(
|
||||
category=Constants.OUTPUT,
|
||||
payload=StepOutput.build_finished_output(name=f"{message.caller or self.name()}",
|
||||
step_num=0,
|
||||
task_id=self.task_id),
|
||||
sender=self.name(),
|
||||
receiver=agents[0].tool_name,
|
||||
session_id=session_id,
|
||||
headers=headers
|
||||
)
|
||||
|
||||
for agent in agents:
|
||||
async for event in self._agent(agent, message):
|
||||
logger.info(f"agent handler send message: {event}")
|
||||
yield event
|
||||
|
||||
async def _agent(self, action: ActionModel, message: Message):
|
||||
self.agent_calls.append(action.agent_name)
|
||||
agent = self.swarm.agents.get(action.agent_name)
|
||||
# be handoff
|
||||
agent_name = action.tool_name
|
||||
if not agent_name:
|
||||
async for event in self._stop_check(action, message):
|
||||
yield event
|
||||
return
|
||||
|
||||
headers = {"context": message.context}
|
||||
session_id = message.session_id
|
||||
cur_agent = self.swarm.agents.get(agent_name)
|
||||
if not cur_agent or not agent:
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg=f"Can not find {agent_name} or {action.agent_name} agent in swarm.",
|
||||
data=action,
|
||||
stop=True),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=headers
|
||||
)
|
||||
return
|
||||
|
||||
cur_agent._finished = False
|
||||
con = action.policy_info
|
||||
if action.params and 'content' in action.params:
|
||||
con = action.params['content']
|
||||
observation = Observation(content=con, observer=agent.id(), from_agent_name=agent.id())
|
||||
|
||||
if agent.handoffs and agent_name not in agent.handoffs:
|
||||
if message.caller:
|
||||
message.receiver = message.caller
|
||||
message.caller = ''
|
||||
yield message
|
||||
else:
|
||||
yield Message(category=Constants.TASK,
|
||||
payload=TaskItem(msg=f"Can not handoffs {agent_name} agent ", data=observation),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.RERUN,
|
||||
headers=headers)
|
||||
return
|
||||
|
||||
headers = message.headers.copy()
|
||||
# headers.update({"agent_as_tool": True})
|
||||
yield Message(
|
||||
category=Constants.AGENT,
|
||||
payload=observation,
|
||||
caller=message.caller,
|
||||
sender=action.agent_name,
|
||||
session_id=session_id,
|
||||
receiver=action.tool_name,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
async def _stop_check(self, action: ActionModel, message: Message) -> AsyncGenerator[Message, None]:
|
||||
if GraphBuildType.TEAM.value == self.swarm.build_type:
|
||||
async for event in self._team_stop_check(action, message):
|
||||
yield event
|
||||
elif GraphBuildType.HANDOFF.value == self.swarm.build_type:
|
||||
async for event in self._handoff_stop_check(action, message):
|
||||
yield event
|
||||
else:
|
||||
async for event in self._workflow_stop_check(action, message):
|
||||
yield event
|
||||
|
||||
async def _workflow_stop_check(self, action: ActionModel, message: Message) -> AsyncGenerator[Message, None]:
|
||||
# Equivalent to scheduling
|
||||
session_id = message.session_id
|
||||
agent_name = action.agent_name
|
||||
agent = self.swarm.agents.get(agent_name)
|
||||
if not agent:
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(
|
||||
msg=f"Can not find {action.agent_name} agent in ordered_agents: {self.swarm.ordered_agents}.",
|
||||
data=action,
|
||||
stop=True),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=message.headers
|
||||
)
|
||||
return
|
||||
|
||||
receiver = None
|
||||
# loop agent type
|
||||
if isinstance(agent, LoopableAgent):
|
||||
agent.cur_run_times += 1
|
||||
if not agent.finished:
|
||||
receiver = agent.goto
|
||||
|
||||
if receiver:
|
||||
yield Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content=action.policy_info),
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
receiver=receiver,
|
||||
headers=message.headers
|
||||
)
|
||||
else:
|
||||
agent_graph: AgentGraph = self.swarm.agent_graph
|
||||
# next
|
||||
successor = agent_graph.successor.get(agent_name)
|
||||
if not successor:
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=action.policy_info,
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.FINISHED,
|
||||
headers=message.headers
|
||||
)
|
||||
return
|
||||
|
||||
for k, _ in successor.items():
|
||||
predecessor = agent_graph.predecessor.get(k)
|
||||
if not predecessor:
|
||||
raise AWorldRuntimeException(f"{k} has no predecessor {agent_name}, may changed during iteration.")
|
||||
|
||||
all_input = {}
|
||||
pre_finished = True
|
||||
for pre_k, _ in predecessor.items():
|
||||
if pre_k == agent_name:
|
||||
all_input[agent_name] = action.policy_info
|
||||
continue
|
||||
# check all predecessor agent finished
|
||||
run_node: RunNode = self.runner.state_manager.query_by_task(
|
||||
task_id=message.context.get_task().id,
|
||||
busi_typ=RunNodeBusiType.AGENT,
|
||||
busi_id=pre_k
|
||||
)
|
||||
if run_node:
|
||||
run_node = run_node[0]
|
||||
else:
|
||||
raise AWorldRuntimeException(f"{pre_k} can't find in task: {message.context.get_task().id}.")
|
||||
if run_node.status == RunNodeStatus.RUNNING or run_node.status == RunNodeStatus.INIT:
|
||||
# mean not finished
|
||||
pre_finished = False
|
||||
logger.info(f"{pre_k} not finished, will wait it.")
|
||||
else:
|
||||
logger.info(f"{pre_k} finished, result is: {run_node.results}")
|
||||
payload = run_node.results[-1].result.payload[0]
|
||||
all_input[pre_k] = payload.policy_info
|
||||
|
||||
if pre_finished:
|
||||
yield Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content=all_input if len(all_input) > 1 else all_input.get(agent_name)),
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
receiver=k,
|
||||
headers=message.headers
|
||||
)
|
||||
|
||||
async def _team_stop_check(self, action: ActionModel, message: Message) -> AsyncGenerator[Message, None]:
|
||||
caller = message.caller
|
||||
session_id = message.session_id
|
||||
agent = self.swarm.agents.get(action.agent_name)
|
||||
if ((not caller or caller == self.swarm.communicate_agent.id())
|
||||
and (self.swarm.cur_step >= self.swarm.max_steps or self.swarm.finished or
|
||||
(agent.id() == self.swarm.agent_graph.root_agent.id() and agent.finished))):
|
||||
logger.info(
|
||||
f"FINISHED|_social_stop_check finished|{self.swarm.cur_step}|{self.swarm.max_steps}|{self.swarm.finished}")
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=action.policy_info,
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.FINISHED,
|
||||
headers={"context": message.context}
|
||||
)
|
||||
agent = self.swarm.agents.get(action.agent_name)
|
||||
caller = self.swarm.agent_graph.root_agent.id() or message.caller
|
||||
if agent.id() != self.swarm.agent_graph.root_agent.id():
|
||||
logger.info(f"_stop_check Team|{agent.id()} --> {caller}")
|
||||
yield Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content=action.policy_info),
|
||||
sender=agent.id(),
|
||||
session_id=message.session_id,
|
||||
receiver=caller,
|
||||
headers=message.headers
|
||||
)
|
||||
|
||||
async def _handoff_stop_check(self, action: ActionModel, message: Message) -> AsyncGenerator[Message, None]:
|
||||
headers = {"context": message.context}
|
||||
agent = self.swarm.agents.get(action.agent_name)
|
||||
caller = message.caller
|
||||
session_id = message.session_id
|
||||
if endless_detect(self.agent_calls,
|
||||
endless_threshold=self.endless_threshold,
|
||||
root_agent_name=self.swarm.communicate_agent.id()):
|
||||
logger.info(
|
||||
f"FINISHED|_social_stop_check endless_detect|{self.agent_calls}|{self.endless_threshold}|{self.swarm.communicate_agent.id()}")
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=action.policy_info,
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.FINISHED,
|
||||
headers=headers
|
||||
)
|
||||
return
|
||||
|
||||
if not caller or caller == self.swarm.communicate_agent.id():
|
||||
if self.swarm.cur_step >= self.swarm.max_steps or self.swarm.finished:
|
||||
logger.info(
|
||||
f"FINISHED|_social_stop_check finished|{self.swarm.cur_step}|{self.swarm.max_steps}|{self.swarm.finished}")
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=action.policy_info,
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.FINISHED,
|
||||
headers=headers
|
||||
)
|
||||
else:
|
||||
self.swarm.cur_step += 1
|
||||
logger.info(f"_social_stop_check execute loop {self.swarm.cur_step}.")
|
||||
yield Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content=action.policy_info),
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
receiver=self.swarm.communicate_agent.id(),
|
||||
headers=message.headers
|
||||
)
|
||||
else:
|
||||
idx = 0
|
||||
for idx, name in enumerate(self.agent_calls[::-1]):
|
||||
if name == agent.id():
|
||||
break
|
||||
idx = len(self.agent_calls) - idx - 1
|
||||
if idx:
|
||||
caller = self.agent_calls[idx - 1]
|
||||
|
||||
yield Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content=action.policy_info),
|
||||
sender=agent.id(),
|
||||
session_id=session_id,
|
||||
receiver=caller,
|
||||
headers=message.headers
|
||||
)
|
||||
|
||||
def is_group_finish(self, event: Message) -> bool:
|
||||
"""Determine if an event triggers group completion"""
|
||||
if not isinstance(event, Message) or not event.group_id:
|
||||
return False
|
||||
|
||||
agent_id = event.sender
|
||||
if not agent_id:
|
||||
return False
|
||||
|
||||
agent = self.swarm.agents.get(agent_id)
|
||||
if not agent:
|
||||
return False
|
||||
|
||||
return agent._finished and agent.id() == event.headers.get('root_agent_id', '')
|
||||
|
||||
async def post_handle(self, input: Message, output: Message) -> Message:
|
||||
new_context = output.context.deep_copy()
|
||||
new_context._task = output.context.get_task()
|
||||
output.context = new_context
|
||||
if self.is_group_finish(output):
|
||||
from aworld.runners.state_manager import RuntimeStateManager
|
||||
state_mng = RuntimeStateManager.instance()
|
||||
await state_mng.finish_sub_group(output.group_id, output.headers.get('root_message_id'),
|
||||
[output])
|
||||
return None
|
||||
return output
|
||||
@@ -0,0 +1,102 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
import abc
|
||||
import time
|
||||
|
||||
from typing import TypeVar, Generic, AsyncGenerator
|
||||
|
||||
from aworld.events.util import send_message
|
||||
|
||||
from aworld.core.common import TaskItem
|
||||
|
||||
from aworld.core.event.base import Message, Constants, TopicType, CancelMessage
|
||||
from aworld.logs.util import logger
|
||||
|
||||
IN = TypeVar('IN')
|
||||
OUT = TypeVar('OUT')
|
||||
|
||||
|
||||
class Handler(Generic[IN, OUT]):
|
||||
__metaclass__ = abc.ABCMeta
|
||||
|
||||
@abc.abstractmethod
|
||||
async def handle(self, data: IN) -> AsyncGenerator[OUT, None]:
|
||||
"""Process the data as the expected result.
|
||||
|
||||
Args:
|
||||
data: Data generated while running the task.
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def name(cls):
|
||||
"""Handler name."""
|
||||
return cls.__name__
|
||||
|
||||
|
||||
class DefaultHandler(Handler[Message, AsyncGenerator[Message, None]]):
|
||||
"""Default handler."""
|
||||
|
||||
def __init__(self, runner: 'TaskEventRunner'):
|
||||
self.runner = runner
|
||||
self.hooks = None
|
||||
|
||||
def get_registered_name(self):
|
||||
"""Get the registered name of the handler.
|
||||
|
||||
If the class has a REGISTERED_NAME attribute, return the value of the attribute;
|
||||
otherwise return None.
|
||||
"""
|
||||
return getattr(self.__class__, "REGISTERED_NAME", None)
|
||||
|
||||
def is_valid_message(self, message: Message):
|
||||
"""Validate if the message is valid for this handler.
|
||||
|
||||
If the class has a REGISTERED_NAME attribute, check if the message's category matches the registered name;
|
||||
otherwise return True.
|
||||
"""
|
||||
registered_name = self.get_registered_name()
|
||||
if registered_name is not None:
|
||||
return message.category == registered_name
|
||||
return True
|
||||
|
||||
async def handle(self, message: Message) -> AsyncGenerator[Message, None]:
|
||||
if not self.is_valid_message(message):
|
||||
return
|
||||
timeout = message.context.get_task().timeout
|
||||
time_cost = time.time() - self.runner.start_time
|
||||
if message.topic != TopicType.CANCEL and timeout > 0 and time_cost > timeout:
|
||||
logger.warn(
|
||||
f"[{self.name()}] {message.context.get_task().id} task timeout after {time_cost} seconds.")
|
||||
yield CancelMessage(
|
||||
payload=TaskItem(msg="task timeout.", data=message, stop=True),
|
||||
sender=self.name(),
|
||||
session_id=self.runner.context.session_id,
|
||||
headers={"context": message.context}
|
||||
)
|
||||
return
|
||||
async for event in self._do_handle(message):
|
||||
msg = await self.post_handle(input=message, output=event)
|
||||
if msg:
|
||||
yield msg
|
||||
|
||||
async def _do_handle(self, message: Message) -> AsyncGenerator[Message, None]:
|
||||
yield message
|
||||
|
||||
async def post_handle(self, input:Message, output: Message) -> Message:
|
||||
"""Post handle the message.
|
||||
Args:
|
||||
message: Message generated while running the task.
|
||||
"""
|
||||
return output
|
||||
|
||||
async def run_hooks(self, message: Message, hook_point: str) -> AsyncGenerator[Message, None]:
|
||||
if not self.hooks:
|
||||
return
|
||||
hooks = self.hooks.get(hook_point, [])
|
||||
for hook in hooks:
|
||||
try:
|
||||
msg = await hook.exec(message)
|
||||
if msg:
|
||||
yield msg
|
||||
except:
|
||||
logger.warning(f"{self.name()}|{hook.point()} {hook.name()} execute fail.")
|
||||
@@ -0,0 +1,413 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
import abc
|
||||
import copy
|
||||
import json
|
||||
from typing import AsyncGenerator, List, Dict, Any, Tuple
|
||||
|
||||
from aworld.agents.llm_agent import Agent
|
||||
from aworld.core.agent.base import is_agent
|
||||
from aworld.core.common import ActionModel, TaskItem, Observation, ActionResult
|
||||
from aworld.core.context.base import Context
|
||||
from aworld.core.event.base import Message, Constants, TopicType, GroupMessage
|
||||
from aworld.logs.util import logger
|
||||
from aworld.output.base import StepOutput
|
||||
from aworld.runners import HandlerFactory
|
||||
from aworld.runners.handler.base import DefaultHandler
|
||||
from aworld.runners.handler.tool import DefaultToolHandler
|
||||
from aworld.runners.state_manager import RuntimeStateManager, RunNodeStatus
|
||||
from aworld.utils.serialized_util import to_serializable
|
||||
from aworld.utils.run_util import exec_agent
|
||||
|
||||
|
||||
class GroupHandler(DefaultHandler):
|
||||
__metaclass__ = abc.ABCMeta
|
||||
|
||||
def __init__(self, runner: 'TaskEventRunner'):
|
||||
super().__init__(runner)
|
||||
self.runner = runner
|
||||
self.swarm = runner.swarm
|
||||
self.endless_threshold = runner.endless_threshold
|
||||
self.task_id = runner.task.id
|
||||
|
||||
@classmethod
|
||||
def name(cls):
|
||||
return "_group_handler"
|
||||
|
||||
|
||||
@HandlerFactory.register(name=f'__{Constants.GROUP}__')
|
||||
class DefaultGroupHandler(GroupHandler):
|
||||
def is_valid_message(self, message: Message):
|
||||
if message.category != Constants.GROUP:
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _do_handle(self, message: GroupMessage) -> AsyncGenerator[Message, None]:
|
||||
if not self.is_valid_message(message):
|
||||
return
|
||||
|
||||
self.context = message.context
|
||||
group_id = message.group_id
|
||||
headers = {'context': self.context}
|
||||
state_manager = RuntimeStateManager.instance()
|
||||
if message.topic == TopicType.GROUP_ACTIONS:
|
||||
# message.payload is List[ActionModel]
|
||||
node_ids = []
|
||||
action_messages = []
|
||||
agents = []
|
||||
tools = []
|
||||
agent_actions_map = {}
|
||||
for action in message.payload:
|
||||
if not isinstance(action, ActionModel):
|
||||
# error message, p2p
|
||||
async for event in self._send_failed_message(message, message.payload, message):
|
||||
yield event
|
||||
return
|
||||
if is_agent(action):
|
||||
agents.append(action)
|
||||
agent_name = action.tool_name
|
||||
if agent_name not in agent_actions_map:
|
||||
agent_actions_map[agent_name] = []
|
||||
agent_actions_map[agent_name].append(action)
|
||||
else:
|
||||
tools.append(action)
|
||||
|
||||
# Process each agent's actions
|
||||
agent_messages = {}
|
||||
for agent_name, actions in agent_actions_map.items():
|
||||
# Get original agent
|
||||
original_agent = self.swarm.agents.get(agent_name)
|
||||
if not original_agent:
|
||||
error_msg = Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg=f"Can not find {agent_name} agent in swarm.",
|
||||
data=actions,
|
||||
stop=True),
|
||||
sender=self.name(),
|
||||
session_id=message.session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers={'context': self.context}
|
||||
)
|
||||
yield error_msg
|
||||
return
|
||||
|
||||
# Create agent copies and execute for each action
|
||||
for action in actions:
|
||||
msg = await self._build_agent_message(action, message)
|
||||
if msg.category != Constants.AGENT:
|
||||
yield msg
|
||||
return
|
||||
self._update_headers(msg, message)
|
||||
agent_copy = self.copy_agent(original_agent)
|
||||
con = action.policy_info
|
||||
if action.params and 'content' in action.params:
|
||||
con = action.params['content']
|
||||
|
||||
if agent_name not in agent_messages:
|
||||
agent_messages[agent_name] = []
|
||||
agent_messages[agent_name].append((con, agent_copy, msg))
|
||||
agent_node_ids, agent_tasks = await self._parallel_exec_agents_actions(agent_messages, message)
|
||||
node_ids.extend(agent_node_ids)
|
||||
|
||||
if tools:
|
||||
tool_mapping = {}
|
||||
for action in tools:
|
||||
tool_name = action.tool_name
|
||||
if tool_name not in tool_mapping:
|
||||
tool_mapping[tool_name] = []
|
||||
tool_mapping[tool_name].append(action)
|
||||
for tool_name, actions in tool_mapping.items():
|
||||
msg = await self._build_tool_message(actions, message)
|
||||
self._update_headers(msg, message)
|
||||
action_messages.append(msg)
|
||||
node_ids.append(msg.id)
|
||||
|
||||
# create group
|
||||
group_meta_data = message.headers.copy()
|
||||
group_meta_data["context"] = message.context.deep_copy()
|
||||
group_meta_data["context"].set_task(message.context.get_task())
|
||||
await state_manager.create_group(group_id, message.session_id, node_ids,
|
||||
message.headers.get('parent_group_id'),
|
||||
group_meta_data)
|
||||
for _, acts in agent_messages.items():
|
||||
for act in acts:
|
||||
self.runner.state_manager.start_message_node(act[2])
|
||||
for msg in action_messages:
|
||||
yield msg
|
||||
await self.process_agent_tasks(agent_tasks, message)
|
||||
|
||||
elif message.topic == TopicType.GROUP_RESULTS:
|
||||
# merge group results
|
||||
action_results = []
|
||||
group_results = message.payload
|
||||
group_sender = None
|
||||
group_sender_node_id = None
|
||||
agent_context = self.context.deep_copy()
|
||||
agent_context._task = self.context.get_task()
|
||||
receiver_results = {}
|
||||
|
||||
for node_id, handle_res_list in group_results.items():
|
||||
if not handle_res_list:
|
||||
logger.warn(f"{self.name()} get group result with empty handle_res.")
|
||||
return
|
||||
node = state_manager._find_node(node_id)
|
||||
tool_call_id = node.metadata.get('root_tool_call_id')
|
||||
is_tool = not tool_call_id and not node.metadata.get('root_agent_id')
|
||||
if not group_sender:
|
||||
group_sender = node.metadata.get('group_sender')
|
||||
if not group_sender_node_id:
|
||||
group_sender_node_id = node.metadata.get('group_sender_node_id')
|
||||
node_results = []
|
||||
for handle_res in handle_res_list:
|
||||
res_msg = handle_res.result
|
||||
res_status = handle_res.status
|
||||
if res_status == RunNodeStatus.FAILED or not res_msg:
|
||||
logger.warn(f"{self.name()} get group result with failed handle_res: {handle_res}.")
|
||||
return
|
||||
receiver = res_msg.receiver
|
||||
if not receiver:
|
||||
logger.warn(f"{self.name()} get group result with empty receiver: {res_msg}.")
|
||||
continue
|
||||
|
||||
if receiver != group_sender:
|
||||
if receiver not in receiver_results:
|
||||
receiver_results[receiver] = []
|
||||
receiver_results[receiver].append(res_msg)
|
||||
else:
|
||||
if is_tool and isinstance(res_msg.payload, Observation):
|
||||
action_results.extend(res_msg.payload.action_result)
|
||||
else:
|
||||
node_results.append(res_msg.payload)
|
||||
self._merge_context(agent_context, res_msg.context)
|
||||
|
||||
if node_results and tool_call_id:
|
||||
act_res = ActionResult(
|
||||
content=json.dumps(to_serializable(node_results), ensure_ascii=False),
|
||||
tool_call_id=tool_call_id
|
||||
)
|
||||
action_results.append(act_res)
|
||||
if action_results:
|
||||
group_res_msg = Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content="", action_result=action_results),
|
||||
caller=message.caller,
|
||||
sender=self.name(),
|
||||
session_id=message.session_id,
|
||||
receiver=group_sender,
|
||||
headers={'context': agent_context}
|
||||
)
|
||||
receiver_results[group_sender] = [group_res_msg]
|
||||
|
||||
for receiver, res_msgs in receiver_results.items():
|
||||
result_message = self._merge_result_messages(res_msgs, message, group_sender_node_id)
|
||||
group_headers = {}
|
||||
group_sender_node = state_manager._find_node(group_sender_node_id)
|
||||
if group_sender_node:
|
||||
group_headers.update(group_sender_node.metadata.copy())
|
||||
group_headers['level'] = headers.get('level', 0) + 1
|
||||
group_headers['context'] = result_message.context or self.context
|
||||
result_message.headers = group_headers
|
||||
yield result_message
|
||||
|
||||
def copy_agent(self, agent: Agent):
|
||||
"""Create a copy of the agent
|
||||
|
||||
Args:
|
||||
agent: Original agent object
|
||||
|
||||
Returns:
|
||||
Deep copy of the agent
|
||||
"""
|
||||
return agent
|
||||
|
||||
async def _parallel_exec_agents_actions(self, agent_messages: Dict[str, List[Tuple[str, Agent, Message]]],
|
||||
message: Message):
|
||||
"""Execute multiple agent actions in parallel
|
||||
|
||||
Args:
|
||||
agent_messages: Messages for agent actions
|
||||
"""
|
||||
tasks = {}
|
||||
messages_ids = []
|
||||
for agent_name, acts in agent_messages.items():
|
||||
for act in acts:
|
||||
agent_message = act[2]
|
||||
messages_ids.append(agent_message.id)
|
||||
tasks[agent_message.id] = exec_agent(act[0], act[1], self.context, sub_task=True, outputs=self.context.outputs)
|
||||
|
||||
return messages_ids, tasks
|
||||
|
||||
async def process_agent_tasks(self, agent_tasks, input_message):
|
||||
"""Process agent async tasks
|
||||
|
||||
Args:
|
||||
agent_tasks: Agent async tasks
|
||||
"""
|
||||
root_agent_set = set()
|
||||
for node_id, task in agent_tasks.items():
|
||||
res = await task
|
||||
logger.info(f"{node_id} finished task: {res}")
|
||||
state_manager = self.runner.state_manager
|
||||
node = state_manager._find_node(node_id)
|
||||
if not node:
|
||||
logger.warn(f"{self.name()} get group result with empty node.")
|
||||
return
|
||||
root_agent_id = node.metadata.get('root_agent_id')
|
||||
root_agent_set.add(root_agent_id)
|
||||
self.context.merge_sub_context(res.context)
|
||||
msg = Message(
|
||||
category=Constants.AGENT,
|
||||
payload=[ActionModel(policy_info=res.answer, agent_name=root_agent_id)],
|
||||
sender=root_agent_id,
|
||||
session_id=node.session_id,
|
||||
headers={'context': self.context,
|
||||
'root_agent_id': root_agent_id,
|
||||
'root_tool_call_id': node.metadata.get('root_tool_call_id')}
|
||||
)
|
||||
finish_group_messages = []
|
||||
async for event in self.runner._inner_handler_process(
|
||||
results=[msg],
|
||||
handlers=self.runner.handlers
|
||||
):
|
||||
# Only AGENT and TASK messages
|
||||
if isinstance(event, Message) and (
|
||||
event.category == Constants.AGENT or event.category == Constants.TASK):
|
||||
finish_group_messages.append(event)
|
||||
print(f"======== event context: {event.context},.context.task: {event.context.get_task()}")
|
||||
await state_manager.finish_sub_group(node.metadata.get('group_id'), node_id, finish_group_messages)
|
||||
for agent_id in root_agent_set:
|
||||
agent = self.swarm.agents.get(agent_id)
|
||||
if agent:
|
||||
agent._finished = True
|
||||
|
||||
def _merge_result_messages(self, res_msgs: List[Message], input_message: Message, group_sender_node_id: str):
|
||||
"""Merge multiple result messages
|
||||
|
||||
Args:
|
||||
res_msgs: Result messages
|
||||
"""
|
||||
if len(res_msgs) == 1:
|
||||
return res_msgs[0]
|
||||
input_list = []
|
||||
new_context = input_message.context.deep_copy()
|
||||
new_context._task = self.context.get_task()
|
||||
for message in res_msgs:
|
||||
map = {}
|
||||
map[message.sender] = message.payload
|
||||
input_list.append(map)
|
||||
new_context.merge_context(message.context)
|
||||
return Message(
|
||||
category=Constants.AGENT,
|
||||
payload=Observation(content=input_list),
|
||||
sender=self.name(),
|
||||
receiver=res_msgs[0].receiver,
|
||||
session_id=res_msgs[0].session_id,
|
||||
headers={
|
||||
'context': new_context
|
||||
}
|
||||
)
|
||||
|
||||
async def _build_agent_message(self, action: ActionModel, message: Message) -> Message:
|
||||
session_id = message.session_id
|
||||
headers = {
|
||||
"context": message.context,
|
||||
"root_tool_call_id": action.tool_call_id
|
||||
}
|
||||
from_agent = self.swarm.agents.get(action.agent_name)
|
||||
tool_name = action.tool_name
|
||||
|
||||
if not tool_name:
|
||||
logger.warn(f"{self.name()} get agent action with empty tool_name.")
|
||||
return Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg=f"Empty tool_name in group_action: {action}.",
|
||||
data=action,
|
||||
stop=True),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=headers
|
||||
)
|
||||
|
||||
cur_agent = self.swarm.agents.get(tool_name)
|
||||
if not cur_agent:
|
||||
return Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg=f"Can not find {tool_name} agent in swarm.",
|
||||
data=action,
|
||||
stop=True),
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=headers
|
||||
)
|
||||
|
||||
cur_agent._finished = False
|
||||
con = action.policy_info
|
||||
if action.params and 'content' in action.params:
|
||||
con = action.params['content']
|
||||
observation = Observation(content=con, observer=from_agent.id(), from_agent_name=from_agent.id())
|
||||
|
||||
return Message(
|
||||
category=Constants.AGENT,
|
||||
payload=observation,
|
||||
caller=message.caller,
|
||||
sender=action.agent_name,
|
||||
session_id=session_id,
|
||||
receiver=cur_agent.id(),
|
||||
headers=headers
|
||||
)
|
||||
|
||||
async def _build_tool_message(self, actions: List[ActionModel], message: Message):
|
||||
session_id = message.session_id
|
||||
headers = {"context": message.context.deep_copy()}
|
||||
return Message(
|
||||
category=Constants.TOOL,
|
||||
payload=actions,
|
||||
sender=self.name(),
|
||||
session_id=session_id,
|
||||
receiver=DefaultToolHandler.name(),
|
||||
headers=headers
|
||||
)
|
||||
|
||||
async def _send_failed_message(self, message, data, result_msg):
|
||||
yield Message(
|
||||
category=Constants.OUTPUT,
|
||||
payload=StepOutput.build_failed_output(name=f"{message.caller or self.name()}",
|
||||
step_num=0,
|
||||
data=result_msg,
|
||||
task_id=self.task_id),
|
||||
sender=self.name(),
|
||||
session_id=self.context.session_id,
|
||||
headers=message.headers
|
||||
)
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg=result_msg, data=data, stop=True),
|
||||
sender=self.name(),
|
||||
session_id=self.context.session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=message.headers
|
||||
)
|
||||
|
||||
def _update_headers(self, message: Message, parent_message: Message):
|
||||
headers = message.headers.copy()
|
||||
context = message.context.deep_copy()
|
||||
context.set_task(self.context.get_task())
|
||||
headers['context'] = context
|
||||
headers['group_id'] = parent_message.group_id
|
||||
headers['root_message_id'] = message.id
|
||||
headers['root_agent_id'] = message.receiver if message.category == Constants.AGENT else ''
|
||||
headers['level'] = 0
|
||||
headers['group_sender'] = parent_message.sender
|
||||
headers['group_sender_node_id'] = parent_message.id
|
||||
headers['parent_group_id'] = parent_message.headers.get('parent_group_id')
|
||||
message.headers = headers
|
||||
|
||||
def _merge_context(self, context: Context, new_context: Context):
|
||||
if not new_context:
|
||||
return
|
||||
if not context:
|
||||
context = new_context
|
||||
return
|
||||
context.merge_context(new_context)
|
||||
@@ -0,0 +1,95 @@
|
||||
# aworld/runners/handler/output.py
|
||||
import json
|
||||
from typing import AsyncGenerator
|
||||
from aworld.core.task import TaskResponse
|
||||
from aworld.models.model_response import ModelResponse
|
||||
from aworld.runners import HandlerFactory
|
||||
from aworld.runners.handler.base import DefaultHandler
|
||||
from aworld.output.base import StepOutput, MessageOutput, Output
|
||||
from aworld.core.common import TaskItem
|
||||
from aworld.core.event.base import Message, Constants, TopicType
|
||||
from aworld.logs.util import logger
|
||||
from aworld.runners.hook.hook_factory import HookFactory
|
||||
from aworld.runners.hook.hooks import HookPoint
|
||||
|
||||
|
||||
@HandlerFactory.register(name=f'__{Constants.OUTPUT}__')
|
||||
class DefaultOutputHandler(DefaultHandler):
|
||||
def __init__(self, runner):
|
||||
super().__init__(runner)
|
||||
self.runner = runner
|
||||
self.hooks = {}
|
||||
if runner.task.hooks:
|
||||
for k, vals in runner.task.hooks.items():
|
||||
self.hooks[k] = []
|
||||
for v in vals:
|
||||
cls = HookFactory.get_class(v)
|
||||
if cls:
|
||||
self.hooks[k].append(cls)
|
||||
|
||||
def is_valid_message(self, message: Message):
|
||||
if message.category != Constants.OUTPUT:
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _do_handle(self, message):
|
||||
if not self.is_valid_message(message):
|
||||
return
|
||||
# 1. get outputs
|
||||
outputs = self.runner.task.outputs
|
||||
if not outputs:
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg="Cannot get outputs.",
|
||||
data=message, stop=True),
|
||||
sender=self.name(),
|
||||
session_id=self.runner.context.session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers={"context": message.context}
|
||||
)
|
||||
return
|
||||
|
||||
# 2. Call OUTPUT_PROCESS hooks to process data in the message
|
||||
async for event in self.run_hooks(message, HookPoint.OUTPUT_PROCESS):
|
||||
# If hook returns a processed message, use the processed message
|
||||
if event and isinstance(event, Message) and event.payload:
|
||||
message.payload = event.payload
|
||||
|
||||
# 3. build Output
|
||||
payload = message.payload
|
||||
mark_complete = False
|
||||
output = None
|
||||
try:
|
||||
if isinstance(payload, Output):
|
||||
output = payload
|
||||
output.task_id = self.runner.task.id
|
||||
elif isinstance(payload, TaskResponse):
|
||||
logger.info(
|
||||
f"FINISHED|output get task_response with usage: {json.dumps(payload.usage)}")
|
||||
if message.topic == TopicType.FINISHED or message.topic == TopicType.ERROR:
|
||||
mark_complete = True
|
||||
elif isinstance(payload, ModelResponse) or isinstance(payload, AsyncGenerator):
|
||||
output = MessageOutput(source=payload, task_id=self.runner.task.id)
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to parse output: {e}")
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg="Failed to parse output.",
|
||||
data=payload, stop=True),
|
||||
sender=self.name(),
|
||||
session_id=self.runner.context.session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers={"context": message.context}
|
||||
)
|
||||
finally:
|
||||
if output:
|
||||
if not output.metadata:
|
||||
output.metadata = {}
|
||||
output.metadata['sender'] = message.sender
|
||||
output.metadata['receiver'] = message.receiver
|
||||
await outputs.add_output(output)
|
||||
if mark_complete:
|
||||
logger.info(f"FINISHED|output mark_completed|{self.runner.task.id}")
|
||||
await outputs.mark_completed()
|
||||
|
||||
return
|
||||
@@ -0,0 +1,138 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
import abc
|
||||
import time
|
||||
|
||||
from typing import AsyncGenerator, TYPE_CHECKING
|
||||
|
||||
from aworld.core.common import TaskItem
|
||||
from aworld.core.tool.base import Tool, AsyncTool
|
||||
|
||||
from aworld.core.event.base import Message, Constants, TopicType
|
||||
from aworld.core.task import TaskResponse
|
||||
from aworld.logs.util import logger
|
||||
from aworld.output import Output
|
||||
from aworld.runners import HandlerFactory
|
||||
from aworld.runners.handler.base import DefaultHandler
|
||||
from aworld.runners.hook.hook_factory import HookFactory
|
||||
from aworld.runners.hook.hooks import HookPoint
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from aworld.runners.event_runner import TaskEventRunner
|
||||
|
||||
|
||||
class TaskHandler(DefaultHandler):
|
||||
__metaclass__ = abc.ABCMeta
|
||||
|
||||
def __init__(self, runner: 'TaskEventRunner'):
|
||||
super().__init__(runner)
|
||||
self.runner = runner
|
||||
self.retry_count = runner.task.max_retry_count
|
||||
self.hooks = {}
|
||||
if runner.task.hooks:
|
||||
for k, vals in runner.task.hooks.items():
|
||||
self.hooks[k] = []
|
||||
for v in vals:
|
||||
cls = HookFactory.get_class(v)
|
||||
if cls:
|
||||
self.hooks[k].append(cls)
|
||||
|
||||
@classmethod
|
||||
def name(cls):
|
||||
return "_task_handler"
|
||||
|
||||
|
||||
@HandlerFactory.register(name=f'__{Constants.TASK}__')
|
||||
class DefaultTaskHandler(TaskHandler):
|
||||
def is_valid_message(self, message: Message):
|
||||
if message.category != Constants.TASK:
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _do_handle(self, message: Message) -> AsyncGenerator[Message, None]:
|
||||
if not self.is_valid_message(message):
|
||||
return
|
||||
|
||||
logger.debug(f"task handler receive message: {message}")
|
||||
|
||||
headers = {"context": message.context}
|
||||
topic = message.topic
|
||||
task_item: TaskItem = message.payload
|
||||
if topic == TopicType.SUBSCRIBE_TOOL:
|
||||
new_tools = message.payload.data
|
||||
for name, tool in new_tools.items():
|
||||
if isinstance(tool, Tool) or isinstance(tool, AsyncTool):
|
||||
await self.runner.event_mng.register(Constants.TOOL, name, tool.step)
|
||||
logger.info(f"dynamic register {name} tool.")
|
||||
else:
|
||||
logger.warning(f"Unknown tool instance: {tool}")
|
||||
return
|
||||
elif topic == TopicType.SUBSCRIBE_AGENT:
|
||||
return
|
||||
elif topic == TopicType.ERROR:
|
||||
async for event in self.run_hooks(message, HookPoint.ERROR):
|
||||
yield event
|
||||
|
||||
logger.warning(f"task {self.runner.task.id} stop, cause: {task_item.msg}")
|
||||
self.runner._task_response = TaskResponse(msg=task_item.msg,
|
||||
answer='',
|
||||
context=message.context,
|
||||
success=False,
|
||||
id=self.runner.task.id,
|
||||
time_cost=(time.time() - self.runner.start_time),
|
||||
usage=self.runner.context.token_usage)
|
||||
if not self.runner.task.is_sub_task:
|
||||
logger.info(f"FINISHED|DefaultTaskHandler|outputs|{self.runner.task.id} {self.runner.task.is_sub_task}")
|
||||
await self.runner.task.outputs.mark_completed()
|
||||
await self.runner.stop()
|
||||
elif topic == TopicType.FINISHED:
|
||||
async for event in self.run_hooks(message, HookPoint.FINISHED):
|
||||
yield event
|
||||
|
||||
self.runner._task_response = TaskResponse(answer=message.payload,
|
||||
success=True,
|
||||
context=message.context,
|
||||
id=self.runner.task.id,
|
||||
time_cost=(time.time() - self.runner.start_time),
|
||||
usage=self.runner.context.token_usage)
|
||||
|
||||
logger.info(f"FINISHED|task|{self.runner.task.id} finished. {self.runner.task.is_sub_task}")
|
||||
if not self.runner.task.is_sub_task:
|
||||
logger.info(f"FINISHED|DefaultTaskHandler|outputs|{self.runner.task.id} {self.runner.task.is_sub_task}")
|
||||
await self.runner.task.outputs.mark_completed()
|
||||
await self.runner.stop()
|
||||
elif topic == TopicType.START:
|
||||
async for event in self.run_hooks(message, HookPoint.START):
|
||||
yield event
|
||||
|
||||
logger.info(f"task start event: {message}, will send init message.")
|
||||
if message.payload:
|
||||
yield message
|
||||
else:
|
||||
yield self.runner.init_message
|
||||
elif topic == TopicType.OUTPUT:
|
||||
yield message
|
||||
elif topic == TopicType.HUMAN_CONFIRM:
|
||||
logger.warn("=============== Get human confirm, pause execution ===============")
|
||||
if self.runner.task.outputs and message.payload:
|
||||
await self.runner.task.outputs.add_output(Output(data=message.payload))
|
||||
self.runner._task_response = TaskResponse(answer=message.payload,
|
||||
success=True,
|
||||
context=message.context,
|
||||
id=self.runner.task.id,
|
||||
time_cost=(time.time() - self.runner.start_time),
|
||||
usage=self.runner.context.token_usage)
|
||||
await self.runner.stop()
|
||||
elif topic == TopicType.CANCEL:
|
||||
# Avoid waiting to receive events and send a mock event for quick cancel
|
||||
yield Message(session_id=self.runner.context.session_id, sender=self.name(), category='mock', headers={"context": message.context})
|
||||
# mark task response as cancelled
|
||||
self.runner._task_response = TaskResponse(answer='',
|
||||
success=False,
|
||||
context=message.context,
|
||||
id=self.runner.task.id,
|
||||
time_cost=(time.time() - self.runner.start_time),
|
||||
usage=self.runner.context.token_usage,
|
||||
msg=f'cancellation message received: {task_item.msg}',
|
||||
status='cancelled')
|
||||
await self.runner.stop()
|
||||
@@ -0,0 +1,123 @@
|
||||
# coding: utf-8
|
||||
# Copyright (c) 2025 inclusionAI.
|
||||
import abc
|
||||
from typing import AsyncGenerator
|
||||
|
||||
from aworld.config import ConfigDict
|
||||
from aworld.core.agent.base import is_agent
|
||||
from aworld.core.common import ActionModel, TaskItem
|
||||
from aworld.core.event.base import Message, Constants, TopicType
|
||||
from aworld.core.tool.base import AsyncTool, Tool, ToolFactory
|
||||
from aworld.logs.util import logger
|
||||
from aworld.runners import HandlerFactory
|
||||
from aworld.runners.handler.base import DefaultHandler
|
||||
|
||||
|
||||
class ToolHandler(DefaultHandler):
|
||||
__metaclass__ = abc.ABCMeta
|
||||
|
||||
def __init__(self, runner: 'TaskEventRunner'):
|
||||
super().__init__(runner)
|
||||
self.tools = runner.tools
|
||||
self.tools_conf = runner.tools_conf
|
||||
|
||||
@classmethod
|
||||
def name(cls):
|
||||
return "_tool_handler"
|
||||
|
||||
|
||||
@HandlerFactory.register(name=f'__{Constants.TOOL}__')
|
||||
class DefaultToolHandler(ToolHandler):
|
||||
def is_valid_message(self, message: Message):
|
||||
if message.category != Constants.TOOL:
|
||||
return False
|
||||
return True
|
||||
|
||||
async def _do_handle(self, message: Message) -> AsyncGenerator[Message, None]:
|
||||
if not self.is_valid_message(message):
|
||||
return
|
||||
|
||||
headers = {"context": message.context}
|
||||
# data is List[ActionModel]
|
||||
data = message.payload
|
||||
if not data:
|
||||
# error message, p2p
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg="no data to process.", data=data, stop=True),
|
||||
sender='agent_handler',
|
||||
session_id=message.session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=headers
|
||||
)
|
||||
return
|
||||
|
||||
for action in data:
|
||||
if not isinstance(action, ActionModel):
|
||||
# error message, p2p
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(msg="action not a ActionModel.", data=data, stop=True),
|
||||
sender=self.name(),
|
||||
session_id=message.session_id,
|
||||
topic=TopicType.ERROR,
|
||||
headers=headers
|
||||
)
|
||||
return
|
||||
|
||||
new_tools = dict()
|
||||
tool_mapping = dict()
|
||||
# Directly use or use tools after creation.
|
||||
for act in data:
|
||||
if is_agent(act):
|
||||
logger.warning(f"somethings wrong, {act} is an agent.")
|
||||
continue
|
||||
|
||||
if not self.tools or (self.tools and act.tool_name not in self.tools):
|
||||
# dynamic only use default config in module.
|
||||
conf = self.tools_conf.get(act.tool_name)
|
||||
if isinstance(conf, dict):
|
||||
conf = ConfigDict(conf)
|
||||
tool = ToolFactory(act.tool_name, conf=conf, asyn=conf.use_async if conf else False)
|
||||
tool.event_driven = True
|
||||
if isinstance(tool, Tool):
|
||||
tool.reset()
|
||||
elif isinstance(tool, AsyncTool):
|
||||
await tool.reset()
|
||||
tool_mapping[act.tool_name] = []
|
||||
self.tools[act.tool_name] = tool
|
||||
new_tools[act.tool_name] = tool
|
||||
if act.tool_name not in tool_mapping:
|
||||
tool_mapping[act.tool_name] = []
|
||||
tool_mapping[act.tool_name].append(act)
|
||||
|
||||
if new_tools:
|
||||
yield Message(
|
||||
category=Constants.TASK,
|
||||
payload=TaskItem(data=new_tools),
|
||||
sender=self.name(),
|
||||
session_id=message.session_id,
|
||||
topic=TopicType.SUBSCRIBE_TOOL,
|
||||
headers=headers
|
||||
)
|
||||
|
||||
for tool_name, actions in tool_mapping.items():
|
||||
if not (isinstance(self.tools[tool_name], Tool) or isinstance(self.tools[tool_name], AsyncTool)):
|
||||
logger.warning(f"Unsupported tool type: {self.tools[tool_name]}")
|
||||
continue
|
||||
|
||||
# send to the tool
|
||||
yield Message(
|
||||
category=Constants.TOOL,
|
||||
payload=actions,
|
||||
sender=actions[0].agent_name if actions else '',
|
||||
session_id=message.session_id,
|
||||
receiver=tool_name,
|
||||
headers=message.headers
|
||||
)
|
||||
|
||||
async def post_handle(self, input:Message, output: Message) -> Message:
|
||||
new_context = output.context.deep_copy()
|
||||
new_context._task = output.context.get_task()
|
||||
output.context = new_context
|
||||
return output
|
||||
Reference in New Issue
Block a user