diff --git a/app/controller/add_langgraph_chat_route.py b/app/controller/add_langgraph_chat_route.py index 14ea4f1..289f910 100644 --- a/app/controller/add_langgraph_chat_route.py +++ b/app/controller/add_langgraph_chat_route.py @@ -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( - model=create_llm("huoshan-think-pro"), - tools=tool_list, - checkpointer=await PostgresCheckpointerManager.get_instance(), - prompt=""" + 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=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')