feat: langgraph/stream接口支持设置系统提示词
This commit is contained in:
@@ -32,23 +32,19 @@ class ChatParam(BaseModel):
|
||||
|
||||
|
||||
class ChatAgent:
|
||||
agent: Union[CompiledStateGraph, None] = None
|
||||
|
||||
@staticmethod
|
||||
async def get_agent() -> CompiledStateGraph:
|
||||
if not ChatAgent.agent:
|
||||
ChatAgent.agent = create_react_agent(
|
||||
async def get_agent(system_prompt: str | None = None) -> CompiledStateGraph:
|
||||
return create_react_agent(
|
||||
model=create_llm("huoshan-think-pro"),
|
||||
tools=tool_list,
|
||||
checkpointer=await PostgresCheckpointerManager.get_instance(),
|
||||
prompt="""
|
||||
prompt=system_prompt or """
|
||||
- 你是一名擅长使用工具的智能助手,你需要根据用户问题来进行回答,请使用中文进行回答。
|
||||
- 当用户问题需要调用工具时再调用工具,否则按照你的知识来回答问题。
|
||||
- 你每次只能调用一个工具
|
||||
- 特别注意,特别注意,特别注意,如果工具的返回结果是一个数组,你只能回复“工具已经执行完毕”
|
||||
"""
|
||||
)
|
||||
return ChatAgent.agent
|
||||
|
||||
@staticmethod
|
||||
async def get_chat_state(thread_id: str):
|
||||
@@ -96,7 +92,12 @@ def add_langgraph_chat_route(app: FastAPI):
|
||||
|
||||
async def generator_function():
|
||||
|
||||
graph = await ChatAgent.get_agent()
|
||||
system_prompt = stream_input.get('system_prompt', None)
|
||||
print("system_prompt ==>> ", system_prompt)
|
||||
|
||||
graph = await ChatAgent.get_agent(
|
||||
system_prompt=system_prompt
|
||||
)
|
||||
|
||||
# chat_state = await ChatAgent.get_chat_state(thread_id)
|
||||
# chat_history_list = chat_state.get('messages')
|
||||
|
||||
Reference in New Issue
Block a user