update
This commit is contained in:
@@ -0,0 +1,11 @@
|
||||
import os
|
||||
import time
|
||||
import operator
|
||||
|
||||
from typing import TypedDict, List, Annotated
|
||||
from langchain_openai import ChatOpenAI
|
||||
from langgraph.graph import END, StateGraph
|
||||
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
llm = ChatOpenAI(model="gpt-4o", temperature=0.1)
|
||||
@@ -0,0 +1,148 @@
|
||||
import os
|
||||
import time
|
||||
import operator
|
||||
import difflib
|
||||
|
||||
from typing import TypedDict, List, Annotated
|
||||
from langchain_openai import ChatOpenAI
|
||||
from langgraph.graph import END, StateGraph
|
||||
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
||||
from langchain_core.messages import HumanMessage
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
llm = ChatOpenAI(model="gpt-4o", temperature=0.1, openai_api_key='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')
|
||||
|
||||
MEMBERS = {'天气查询模块': '输入一个城市和日期,查询当日的天气状况',
|
||||
'网络查询模块': '在网上搜索指定的内容',
|
||||
'诗歌创作模块': '根据给定的要求,进行诗歌创作',
|
||||
'结束节点': '完成所有任务或者任务执行失败的情况下进入'}
|
||||
|
||||
|
||||
class AgentState(TypedDict):
|
||||
user_query: str
|
||||
messages: Annotated[List[str], operator.add]
|
||||
plan: str
|
||||
members: dict
|
||||
next_agent: str
|
||||
next_plan: str
|
||||
|
||||
|
||||
def planner(state: AgentState):
|
||||
user_query = state.get('user_query', '')
|
||||
planner_prompt = ChatPromptTemplate.from_messages([
|
||||
('system', '''你是一个团队的管理者,你负责根据用户输入的要求,制定一个计划来完成。你的团队成员和能力如下:
|
||||
1. 天气查询模块:输入一个城市和日期,查询当日的天气状况。
|
||||
2. 网络查询模块:输入一个需要查询的内容,在网上搜索指定的内容。
|
||||
3. 诗歌创作模块:根据给定的要求,进行诗歌创作
|
||||
请你根据用户的需求,制定一个处理计划'''),
|
||||
('user', f'当前的用户需求是:'),
|
||||
MessagesPlaceholder(variable_name="query"),
|
||||
])
|
||||
planner_chain = (planner_prompt | llm)
|
||||
|
||||
result_dict = {'plan': '为了完成这个任务,我们需要按照以下步骤进行:\n\n1. **网络查询模块**:\n - 任务:查询2024年中国GDP最高的城市。\n - 行动:在网上搜索2024年中国GDP最高的城市的信息。\n\n2. **天气查询模块**:\n - 任务:查询该城市今天的天气状况。\n - 行动:使用天气查询模块,输入城市名称和当前日期,获取该城市的天气信息。\n\n3. **诗歌创作模块**:\n - 任务:根据查询到的城市名称和天气状况,创作一首诗。\n - 行动:利用诗歌创作模块,结合城市的特点和天气状况,创作一首诗。\n\n请各模块按照上述计划执行任务,并在完成后汇总信息。'}
|
||||
|
||||
return result_dict
|
||||
# llm_resp = planner_chain.invoke({'query': [('user', user_query)]})
|
||||
# return {'plan': llm_resp.content}
|
||||
|
||||
|
||||
def decision_node(state: AgentState):
|
||||
class NextStep(BaseModel):
|
||||
next_agent: str = Field(description='下一步应该交给哪个模块来完成')
|
||||
action: str = Field(description='下一步应该执行的行动计划是什么')
|
||||
|
||||
members = state.get('members', {})
|
||||
member_keys = list(members.keys())
|
||||
decision_prompt = ChatPromptTemplate.from_messages([
|
||||
('system', f'你负责针对一个给定的任务计划,以及当前的任务执行状态,制定下一步应该由谁来执行什么内容。可以在这些模块中进行选择:{", ".join(member_keys)}。'),
|
||||
('user', f'当前的计划是:'),
|
||||
MessagesPlaceholder(variable_name="plan"),
|
||||
('user', '以下是各个模块之间的消息记录'),
|
||||
MessagesPlaceholder(variable_name="messages")
|
||||
])
|
||||
decision_chain = (decision_prompt | llm.with_structured_output(NextStep))
|
||||
llm_resp = decision_chain.invoke({'plan': [('user', state.get('plan', ''))], 'messages': state.get('messages', [])})
|
||||
next_agent = llm_resp.next_agent
|
||||
action = llm_resp.action
|
||||
return {'messages': [HumanMessage(content=f'接下来由{next_agent}来执行{action}')], 'next_agent': next_agent, 'next_plan': action}
|
||||
|
||||
|
||||
def router_node(state: AgentState):
|
||||
members = state.get('members', {})
|
||||
member_keys = list(members.keys())
|
||||
best_match = difflib.get_close_matches(state.get('next_agent', ''), member_keys, n=1, cutoff=0)
|
||||
if not best_match:
|
||||
raise ValueError(f'Select a wrong agent: {state.get("next_agent")}')
|
||||
return {'next': best_match}
|
||||
|
||||
|
||||
def weather_node(state: AgentState):
|
||||
result = '上海的天气是25℃,下雨'
|
||||
return {'messages': [HumanMessage(content=result)]}
|
||||
|
||||
|
||||
def websearch_node(state: AgentState):
|
||||
result = '2024年中国GDP最高的城市是上海'
|
||||
return {'messages': [HumanMessage(content=result)]}
|
||||
|
||||
|
||||
def composer_node(state: AgentState):
|
||||
composer_prompt = ChatPromptTemplate.from_messages([
|
||||
('system', f'你负责根据用户的要求以诗歌的形式进行创作'),
|
||||
('user', f'当前用户的要求是:'),
|
||||
MessagesPlaceholder(variable_name="request"),
|
||||
])
|
||||
composer_chain = composer_prompt | llm
|
||||
llm_resp = composer_chain.invoke({'request': [('user', state.get('next_plan', ''))]})
|
||||
return {'messages': [HumanMessage(content=llm_resp.content)]}
|
||||
|
||||
|
||||
def end_node(state: AgentState):
|
||||
print('placeholder')
|
||||
|
||||
|
||||
workflow = StateGraph(AgentState)
|
||||
|
||||
workflow.add_node('planner', planner)
|
||||
workflow.add_node('决策节点', decision_node)
|
||||
workflow.add_node('router', router_node)
|
||||
workflow.add_node('天气查询模块', weather_node)
|
||||
workflow.add_node('网络查询模块', websearch_node)
|
||||
workflow.add_node('诗歌创作模块', composer_node)
|
||||
workflow.add_node('结束节点', end_node)
|
||||
|
||||
workflow.add_edge('planner', '决策节点')
|
||||
workflow.add_edge('决策节点', 'router')
|
||||
workflow.add_conditional_edges(
|
||||
'router',
|
||||
lambda x: x['next'],
|
||||
{
|
||||
'天气查询模块': '天气查询模块',
|
||||
'网络查询模块': '网络查询模块',
|
||||
'诗歌创作模块': '诗歌创作模块',
|
||||
'结束节点': '结束节点'
|
||||
}
|
||||
)
|
||||
|
||||
workflow.add_edge('天气查询模块', '决策节点')
|
||||
workflow.add_edge('网络查询模块', '决策节点')
|
||||
workflow.add_edge('诗歌创作模块', '决策节点')
|
||||
workflow.add_edge('结束节点', END)
|
||||
|
||||
workflow.set_entry_point('planner')
|
||||
|
||||
graph = workflow.compile()
|
||||
print(graph.get_graph().draw_ascii())
|
||||
|
||||
enter = {'user_query': "查询2024年中国GDP最高的城市,以这个城市今天的天气,创作一首诗",
|
||||
"members": MEMBERS}
|
||||
for s in graph.stream(enter):
|
||||
if "__end__" not in s:
|
||||
print(s)
|
||||
print('-'*80)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
from tavily import TavilyClient
|
||||
import requests
|
||||
|
||||
|
||||
query = '乔布斯是谁'
|
||||
|
||||
base_url = 'https://api.tavily.com/search'
|
||||
json = {
|
||||
'api_key': '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',
|
||||
'query': query,
|
||||
'search_depth': 'basic',
|
||||
'include_answer': False,
|
||||
'include_images': False,
|
||||
'include_raw_content': False,
|
||||
'max_results': 5,
|
||||
'include_domains': [],
|
||||
'exclude_domains': []
|
||||
}
|
||||
|
||||
response = requests.post(base_url, json=json, verify=False)
|
||||
|
||||
if response.status_code == 200:
|
||||
search_result = response.json()
|
||||
else:
|
||||
raise Exception(f'Error: {response.status_code}: {response.reason}')
|
||||
|
||||
print('done')
|
||||
Reference in New Issue
Block a user