import asyncio
import json
from langchain_mcp_adapters.client import MultiServerMCPClient
from typing import Dict, Any, Optional, List
from langchain_core.messages import ToolMessage, HumanMessage, AIMessage, SystemMessage
from langgraph.checkpoint.memory import MemorySaver
from langgraph.graph import MessagesState
from langgraph.graph import StateGraph
from langgraph.constants import END, START
import asyncio
from langgraph.types import interrupt, Command
from agent.env_utils import ZHIPU_API_KEY
from agent.my_llm import llm

"""
(5)自定义某一个工具中断

"""

# 外网公开mcp服务端的连接配置
mcp_zhipu_websearch_config = {
    "url": "https://open.bigmodel.cn/api/mcp/web_search/sse?Authorization=" + ZHIPU_API_KEY,
    'transport': 'sse'
}
mcp_12306_server_config = {
    "url": "https://mcp.api-inference.modelscope.net/29a158ef90c44b/sse",
    'transport': 'sse'
}
mcp_chart_server_config = {
    "url": "https://mcp.api-inference.modelscope.net/281ac94c0a0949/sse",
    'transport': 'sse'
}

# 创建多服务器客户端
mcp_client = MultiServerMCPClient(
    {
        'chart_mcp': mcp_chart_server_config,
        '12306_mcp': mcp_12306_server_config,
        'zhipu_mcp': mcp_zhipu_websearch_config
    }
)

# 改进原始的 BasicToolNode,添加错误处理
class BasicToolNode:
    def __init__(self, tools: list):
        self.tools_by_name = {tool.name: tool for tool in tools}

    async def __call__(self, state: Dict[str, Any]) -> Dict[str, List[ToolMessage]]:
        if not (messages := state.get('messages')):
            raise ValueError('输入的数据中未找到消息内容')
        message: AIMessage = messages[-1]

        # 选择工具进行中断
        tool_name = message.tool_calls[0]['name'] if message.tool_calls else None
        if tool_name in ['webSearchPro', 'generate_column_chart', 'generate_pie_chart', 'webSearchStd',
                         'webSearchSogou']:
            response = interrupt(
                f"AI大模型尝试调用工具:{tool_name},\n"
                f"请审核并选择,批准(y)或这届给我工具执行的答案"
            )
            if response['answer'] == 'y':
                pass
            else:
                return {'messages': [ToolMessage(
                    content=f"人类终止了该工具调用,理由是:{response['answer']}",
                    name=tool_name,
                    tool_call_id=message.tool_calls[0]['id'],
                )]}


        outputs = await self._execute_tool_calls(message.tool_calls)
        return {'messages': outputs}

    async def _execute_tool_calls(self, tool_calls: List[Dict]) -> List[ToolMessage]:
        async def _invoke_tool(tool_call: Dict) -> ToolMessage:
            try:
                tool = self.tools_by_name.get(tool_call['name'])
                if not tool:
                    raise KeyError(f"未注册的工具:{tool_call['name']}")

                if hasattr(tool, 'ainvoke'):
                    tool_result = await tool.ainvoke(tool_call['args'])
                else:
                    loop = asyncio.get_running_loop()
                    tool_result = await loop.run_in_executor(
                        None, tool.invoke, tool_call['args']
                    )

                return ToolMessage(
                    content=json.dumps(tool_result, ensure_ascii=False),
                    name=tool_call['name'],
                    tool_call_id=tool_call['id'],
                )
            except Exception as e:
                # 添加错误处理
                return ToolMessage(
                    content=json.dumps({"error": str(e)}, ensure_ascii=False),
                    name=tool_call['name'],
                    tool_call_id=tool_call['id'],
                )

        return await asyncio.gather(*[_invoke_tool(tool_call) for tool_call in tool_calls])



class State(MessagesState):
    pass


# 定义路由函数
def route_tools_func(state: State):
    """ 动态路由函数,
    state 有些情况下是列表形式[msg1, msg2, msg3]  (可能性不大)
          有些情况下是字典形式 {'messages': [msg1, msg2, msg3]}
    """
    if isinstance(state, list):
        ai_message = state[-1]  # 取到最后一个消息
    elif messages := state.get('messages', []):
        ai_message = messages[-1]
    else:
        raise ValueError(f"NO messages found in input state to tool_edege:{state}")

    if hasattr(ai_message, 'tool_calls') and len(ai_message.tool_calls) > 0:
        return 'tools'

    return END


# 异步版本的创建函数
async def create_async_graph():
    builder = StateGraph(State)

    # 异步获取工具
    tools = await mcp_client.get_tools()

    llm_with_tools = llm.bind_tools(tools)

    # 异步聊天机器人节点
    async def async_chatbot(state: State):
        response = await llm_with_tools.ainvoke(state['messages'])
        return {'messages': [response]}



    # 添加第一个节点
    builder.add_node('chatbot', async_chatbot)
    # 添加第二个节点
    tool_node = BasicToolNode(tools)
    builder.add_node('tools', tool_node)

    # 添加边
    builder.add_edge(START, 'chatbot')
    builder.add_conditional_edges(
        'chatbot',
        route_tools_func,
        path_map={
            'tools': 'tools',
            END: END
        }
    )

    builder.add_edge('tools', 'chatbot')

    memory = MemorySaver()
    graph = builder.compile(checkpointer=memory)
    return graph


# agent = asyncio.run(create_async_graph())

async def run_graph():
    graph = await  create_async_graph()
    config = {
        'configurable': {
            'thread_id': "zf123"
        }
    }

    def print_message(event, result):
        """格式化输出消息
        event = {'messages':
        [HumanMessage(content="你好"),AIMessage(content="你好!我是AI助手")]
        }
        """
        messages = event.get('messages')
        if messages:
            if isinstance(messages, list):
                message = messages[-1]
            if message.__class__.__name__ == 'AIMessage':
                if message.content:
                    result = message.content
            msg_repr = message.pretty_repr(html=True)
            if len(msg_repr) > 1500:
                msg_repr = msg_repr[:1500] + '...(已截断)'
            print(msg_repr)
        return result

    async def execute_graph(user_input: str) -> str:
        """执行工作量的函数"""
        result = ''
        current_state = graph.get_state(config)
        if current_state.next:
            humman_command = Command(resume={'answer': user_input})
            async  for chunk in graph.astream(humman_command, config, stream_mode='values'):
                result = print_message(chunk, result)
            return result
        else:
            async  for chunk in graph.astream({'messages': ('user', user_input)}, config, stream_mode='values'):
                result = print_message(chunk, result)

                # if chunk.get('__interrupt__',None):
                #     print(chunk['__interrupt__'])

        current_state = graph.get_state(config)
        if current_state.next:  # 出现了工作量中断
            result = current_state.interrupts[0].value

        return result

    # 执行工作流
    while True:
        user_input = input('用户:')
        res = await execute_graph(user_input)
        print('AI回复:', res)


if __name__ == '__main__':
    asyncio.run(run_graph())

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐