feat: comments
This commit is contained in:
@@ -16,6 +16,7 @@ from app.utils.llm_utils import create_llm
|
|||||||
from app.utils.postgres_checkpointer import PostgresCheckpointerManager
|
from app.utils.postgres_checkpointer import PostgresCheckpointerManager
|
||||||
|
|
||||||
|
|
||||||
|
# /*---------------------------------------获取时间工具-------------------------------------------*/
|
||||||
@tool(name_or_callable="获取时间工具", description="一个用于获取当前时间的工具,没有参数")
|
@tool(name_or_callable="获取时间工具", description="一个用于获取当前时间的工具,没有参数")
|
||||||
def tool_get_datetime():
|
def tool_get_datetime():
|
||||||
return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||||
@@ -135,21 +136,6 @@ class ChatAgent:
|
|||||||
)
|
)
|
||||||
return ChatAgent.agent
|
return ChatAgent.agent
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
async def chat(human_message: ChatMessage, thread_id: str):
|
|
||||||
graph = await ChatAgent.get_agent()
|
|
||||||
chat_state = await ChatAgent.get_chat_state(thread_id)
|
|
||||||
chat_history_list = chat_state.get('messages')
|
|
||||||
graph_state = await graph.ainvoke(
|
|
||||||
{"messages": [HumanMessage(content=human_message.content, id=human_message.id)]},
|
|
||||||
config={"configurable": {"thread_id": thread_id}}
|
|
||||||
)
|
|
||||||
return {
|
|
||||||
**graph_state,
|
|
||||||
# 这里之所以要+1,是因为我们认为这次的HumanMessage已经在chat_history_list中了,但是实际上并没有,所以这里要+1
|
|
||||||
"messages": graph_state.get('messages')[len(chat_history_list) + 1:],
|
|
||||||
}
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
async def get_chat_state(thread_id: str):
|
async def get_chat_state(thread_id: str):
|
||||||
graph = await ChatAgent.get_agent()
|
graph = await ChatAgent.get_agent()
|
||||||
@@ -163,11 +149,7 @@ class ChatAgent:
|
|||||||
|
|
||||||
|
|
||||||
def add_langgraph_chat_route(app: FastAPI):
|
def add_langgraph_chat_route(app: FastAPI):
|
||||||
# 聊天接口
|
# 流式对话接口
|
||||||
@app.post("/langgraph/chat")
|
|
||||||
async def langgraph_chat(chat_param: ChatParam):
|
|
||||||
return await ChatAgent.chat(chat_param.human_message, chat_param.thread_id)
|
|
||||||
|
|
||||||
@app.post("/langgraph/stream")
|
@app.post("/langgraph/stream")
|
||||||
async def langgraph_stream(body: dict):
|
async def langgraph_stream(body: dict):
|
||||||
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
|
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
|
||||||
@@ -230,6 +212,7 @@ def add_langgraph_chat_route(app: FastAPI):
|
|||||||
|
|
||||||
return StreamingResponse(generator_function(), media_type="text/event-stream")
|
return StreamingResponse(generator_function(), media_type="text/event-stream")
|
||||||
|
|
||||||
|
# 恢复中断接口
|
||||||
@app.post("/langgraph/chat_resume/{thread_id}")
|
@app.post("/langgraph/chat_resume/{thread_id}")
|
||||||
async def langgraph_chat(body: dict, thread_id: str):
|
async def langgraph_chat(body: dict, thread_id: str):
|
||||||
graph = await ChatAgent.get_agent()
|
graph = await ChatAgent.get_agent()
|
||||||
|
|||||||
Reference in New Issue
Block a user