feat:支持工具调用的模型接口:langgraph/stream
This commit is contained in:
@@ -1,24 +1,52 @@
|
||||
import datetime
|
||||
import json
|
||||
import time
|
||||
from typing import Union, Annotated
|
||||
|
||||
from fastapi import FastAPI
|
||||
from langchain_core.messages import HumanMessage
|
||||
from langchain_core.messages import HumanMessage, AIMessage, ToolMessage
|
||||
from langchain_core.tools import tool
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
from langgraph.prebuilt import create_react_agent
|
||||
from langgraph.types import interrupt, Command
|
||||
from pydantic import BaseModel, Field
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
from app.utils.llm_utils import create_llm
|
||||
from app.utils.postgres_checkpointer import PostgresCheckpointerManager
|
||||
|
||||
|
||||
@tool(name_or_callable="tool_get_datetime", description="一个用于获取当前时间的工具,没有参数")
|
||||
@tool(name_or_callable="获取时间工具", description="一个用于获取当前时间的工具,没有参数")
|
||||
def tool_get_datetime():
|
||||
return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
@tool(
|
||||
name_or_callable="乘法工具",
|
||||
description="一个乘法工具,用于计算两个数字相乘"
|
||||
)
|
||||
def multiple_tool(
|
||||
number1: Annotated[float, "第一个要相乘的数字"],
|
||||
number2: Annotated[float, "第二个要相乘的数字"],
|
||||
) -> float:
|
||||
print("\n invoke multiple_tool", number1, number2, '\n')
|
||||
|
||||
return float(number1) * float(number2)
|
||||
|
||||
|
||||
@tool(
|
||||
name_or_callable="加法工具",
|
||||
description="一个加法工具,用于计算两个数字相加"
|
||||
)
|
||||
def add_tool(
|
||||
number1: Annotated[float, "第一个要相加的数字"],
|
||||
number2: Annotated[float, "第二个要相加的数字"],
|
||||
) -> float:
|
||||
print("\n invoke add_tool", number1, number2, '\n')
|
||||
|
||||
return float(number1) + float(number2)
|
||||
|
||||
|
||||
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
|
||||
def tool_book_hotel(
|
||||
hotel_name: Annotated[str, '酒店名称'],
|
||||
@@ -92,7 +120,12 @@ class ChatAgent:
|
||||
if not ChatAgent.agent:
|
||||
ChatAgent.agent = create_react_agent(
|
||||
model=create_llm(),
|
||||
tools=[tool_book_hotel, tool_get_datetime],
|
||||
tools=[
|
||||
tool_book_hotel,
|
||||
tool_get_datetime,
|
||||
add_tool,
|
||||
multiple_tool,
|
||||
],
|
||||
checkpointer=await PostgresCheckpointerManager.get_instance(),
|
||||
prompt="""
|
||||
你是一名擅长使用工具的智能助手,你需要根据用户问题来进行回答,请使用中文进行回答。
|
||||
@@ -135,6 +168,68 @@ def add_langgraph_chat_route(app: FastAPI):
|
||||
async def langgraph_chat(chat_param: ChatParam):
|
||||
return await ChatAgent.chat(chat_param.human_message, chat_param.thread_id)
|
||||
|
||||
@app.post("/langgraph/stream")
|
||||
async def langgraph_stream(body: dict):
|
||||
print(":::::::::::::::::::::::::langgraph_stream:::::::::::::::::::::")
|
||||
print(body)
|
||||
|
||||
stream_input = body.get('input')
|
||||
stream_config = body.get('config')
|
||||
|
||||
print("stream_input", stream_input)
|
||||
print("stream_config", stream_config)
|
||||
|
||||
async def generator_function():
|
||||
|
||||
graph = await ChatAgent.get_agent()
|
||||
# chat_state = await ChatAgent.get_chat_state(thread_id)
|
||||
# chat_history_list = chat_state.get('messages')
|
||||
async for chunk in graph.astream(
|
||||
stream_input,
|
||||
config=stream_config,
|
||||
stream_mode=['messages', 'updates']
|
||||
):
|
||||
print("chunk-------->>>>>>>>>>", chunk)
|
||||
|
||||
result_template = {
|
||||
"choices": [{"delta": {}, "index": 0}],
|
||||
"created": time.time(),
|
||||
"id": "",
|
||||
"usage": None
|
||||
}
|
||||
|
||||
emit_chunk = {"stream_type": chunk[0], }
|
||||
|
||||
if emit_chunk['stream_type'] == "messages":
|
||||
chunk_message = chunk[1][0]
|
||||
else:
|
||||
for k, v in chunk[1].items():
|
||||
chunk_message = v.get('messages')[0]
|
||||
|
||||
emit_chunk = {
|
||||
**emit_chunk,
|
||||
"msg_id": chunk_message.id,
|
||||
"msg_type": chunk_message.type,
|
||||
"msg_content": chunk_message.content,
|
||||
}
|
||||
|
||||
if isinstance(chunk_message, AIMessage):
|
||||
if chunk_message.tool_calls:
|
||||
emit_chunk["tool_calls"] = chunk_message.tool_calls
|
||||
elif isinstance(chunk_message, ToolMessage):
|
||||
emit_chunk["tool_name"] = chunk_message.name
|
||||
|
||||
result_template["choices"][0]["delta"]['content'] = emit_chunk
|
||||
result_template["choices"][0]["delta"]['role'] = 'assistant'
|
||||
|
||||
result_template['id'] = chunk_message.id
|
||||
result_template['created'] = int(time.time())
|
||||
|
||||
yield f"data: {json.dumps(result_template, ensure_ascii=False)}\n\n"
|
||||
# yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(generator_function(), media_type="text/event-stream")
|
||||
|
||||
@app.post("/langgraph/chat_resume/{thread_id}")
|
||||
async def langgraph_chat(body: dict, thread_id: str):
|
||||
graph = await ChatAgent.get_agent()
|
||||
|
||||
Reference in New Issue
Block a user