feat: 拆分tool文件
This commit is contained in:
@@ -1,117 +1,26 @@
|
||||
import datetime
|
||||
import json
|
||||
import time
|
||||
from typing import Union, Annotated, TypedDict
|
||||
from typing import Union
|
||||
|
||||
import httpx
|
||||
from fastapi import FastAPI, Depends
|
||||
from langchain_core.messages import AIMessage, ToolMessage
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
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 langgraph.types import Command
|
||||
from pydantic import BaseModel, Field
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
from app.config.env import env
|
||||
from app.controller.add_user_route import oauth2_scheme
|
||||
from app.model.ConversationModel import ConversationService
|
||||
from app.tools.tool_add import tool_add
|
||||
from app.tools.tool_book_hotel import tool_book_hotel
|
||||
from app.tools.tool_get_datetime import tool_get_datetime
|
||||
from app.tools.tool_multiple import tool_multiply
|
||||
from app.utils.db_utils import AsyncSessionDep
|
||||
from app.utils.llm_utils import create_llm
|
||||
from app.utils.postgres_checkpointer import PostgresCheckpointerManager, AsyncPostgresSaverDep
|
||||
|
||||
|
||||
# /*---------------------------------------获取时间工具-------------------------------------------*/
|
||||
@tool(
|
||||
name_or_callable="tool_get_datetime",
|
||||
description="一个用于获取当前时间的工具,没有参数"
|
||||
)
|
||||
def tool_get_datetime():
|
||||
return f"当前时间为:{datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')}"
|
||||
|
||||
|
||||
@tool(
|
||||
name_or_callable="tool_multiply",
|
||||
description="一个乘法工具,用于计算两个数字相乘"
|
||||
)
|
||||
def tool_multiply(
|
||||
number1: Annotated[float, "第一个要相乘的数字"],
|
||||
number2: Annotated[float, "第二个要相乘的数字"],
|
||||
) -> float:
|
||||
print("\n invoke tool_multiply", number1, number2, '\n')
|
||||
|
||||
return float(number1) * float(number2)
|
||||
|
||||
|
||||
@tool(
|
||||
name_or_callable="tool_add",
|
||||
description="一个加法工具,用于计算两个数字相加"
|
||||
)
|
||||
def tool_add(
|
||||
number1: Annotated[float, "第一个要相加的数字"],
|
||||
number2: Annotated[float, "第二个要相加的数字"],
|
||||
) -> float:
|
||||
print("\n invoke tool_add", number1, number2, '\n')
|
||||
|
||||
return float(number1) + float(number2)
|
||||
|
||||
|
||||
class BookHotelData(TypedDict):
|
||||
hotel_id: str
|
||||
proj_id: str
|
||||
user_id: str
|
||||
|
||||
|
||||
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
|
||||
async def tool_book_hotel(config: RunnableConfig) -> str:
|
||||
"""
|
||||
触发中断,输入酒店预定表单信息
|
||||
"""
|
||||
book_hotel_data: Union[BookHotelData, str] = interrupt({"title": "请输入酒店预定信息", "formCode": "bookHotel"})
|
||||
|
||||
if book_hotel_data == "N":
|
||||
return f"用户选择取消预定酒店"
|
||||
|
||||
hotel_id = book_hotel_data.get('hotel_id')
|
||||
proj_id = book_hotel_data.get('proj_id')
|
||||
user_id = book_hotel_data.get('user_id')
|
||||
token = config['configurable']['token']
|
||||
|
||||
# print("================ tool_book_hotel ================")
|
||||
# print("hotel_id", hotel_id)
|
||||
# print("proj_id", proj_id)
|
||||
# print("user_id", user_id)
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
url=f"http://localhost:{env.server_port}/book_hotel",
|
||||
headers={
|
||||
"Authorization": f"Bearer {token}"
|
||||
},
|
||||
json={
|
||||
"hotel_id": hotel_id,
|
||||
"user_id": user_id,
|
||||
"proj_id": proj_id,
|
||||
},
|
||||
)
|
||||
response_data = response.json()
|
||||
# print("================ response_data ================")
|
||||
# print(response_data)
|
||||
approve_id = response_data.get('graph_state').get('input_approve_id')
|
||||
output_message = "预定成功," + response_data.get('message')
|
||||
return [
|
||||
output_message,
|
||||
{
|
||||
"component": "Link",
|
||||
"props": {
|
||||
"to": f"/pages/approve/approve-detail?id={approve_id}",
|
||||
"children": f"{output_message},也可以点击这里查看审批进度。"
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
id: str = Field(..., description="消息id")
|
||||
type: str = Field(..., description="消息类型")
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
from typing import Annotated
|
||||
|
||||
from langchain_core.tools import tool
|
||||
|
||||
@tool(
|
||||
name_or_callable="tool_add",
|
||||
description="一个加法工具,用于计算两个数字相加"
|
||||
)
|
||||
def tool_add(
|
||||
number1: Annotated[float, "第一个要相加的数字"],
|
||||
number2: Annotated[float, "第二个要相加的数字"],
|
||||
) -> float:
|
||||
print("\n invoke tool_add", number1, number2, '\n')
|
||||
|
||||
return float(number1) + float(number2)
|
||||
@@ -0,0 +1,63 @@
|
||||
from typing import Union, TypedDict
|
||||
|
||||
import httpx
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langchain_core.tools import tool
|
||||
from langgraph.types import interrupt
|
||||
|
||||
from app.config.env import env
|
||||
|
||||
|
||||
class BookHotelData(TypedDict):
|
||||
hotel_id: str
|
||||
proj_id: str
|
||||
user_id: str
|
||||
|
||||
|
||||
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
|
||||
async def tool_book_hotel(config: RunnableConfig) -> str:
|
||||
"""
|
||||
触发中断,输入酒店预定表单信息
|
||||
"""
|
||||
book_hotel_data: Union[BookHotelData, str] = interrupt({"title": "请输入酒店预定信息", "formCode": "bookHotel"})
|
||||
|
||||
if book_hotel_data == "N":
|
||||
return f"用户选择取消预定酒店"
|
||||
|
||||
hotel_id = book_hotel_data.get('hotel_id')
|
||||
proj_id = book_hotel_data.get('proj_id')
|
||||
user_id = book_hotel_data.get('user_id')
|
||||
token = config['configurable']['token']
|
||||
|
||||
# print("================ tool_book_hotel ================")
|
||||
# print("hotel_id", hotel_id)
|
||||
# print("proj_id", proj_id)
|
||||
# print("user_id", user_id)
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
response = await client.post(
|
||||
url=f"http://localhost:{env.server_port}/book_hotel",
|
||||
headers={
|
||||
"Authorization": f"Bearer {token}"
|
||||
},
|
||||
json={
|
||||
"hotel_id": hotel_id,
|
||||
"user_id": user_id,
|
||||
"proj_id": proj_id,
|
||||
},
|
||||
)
|
||||
response_data = response.json()
|
||||
# print("================ response_data ================")
|
||||
# print(response_data)
|
||||
approve_id = response_data.get('graph_state').get('input_approve_id')
|
||||
output_message = "预定成功," + response_data.get('message')
|
||||
return [
|
||||
output_message,
|
||||
{
|
||||
"component": "Link",
|
||||
"props": {
|
||||
"to": f"/pages/approve/approve-detail?id={approve_id}",
|
||||
"children": f"{output_message},也可以点击这里查看审批进度。"
|
||||
}
|
||||
}
|
||||
]
|
||||
@@ -0,0 +1,31 @@
|
||||
import datetime
|
||||
|
||||
from langchain_core.tools import tool
|
||||
|
||||
|
||||
@tool(name_or_callable="tool_get_datetime", description="一个获取当前时间信息的一个工具")
|
||||
def tool_get_datetime():
|
||||
now = datetime.datetime.now()
|
||||
current_time = now.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
yesterday = (now - datetime.timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S")
|
||||
tomorrow = (now + datetime.timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S")
|
||||
day_after_tomorrow = (now + datetime.timedelta(days=2)).strftime("%Y-%m-%d %H:%M:%S")
|
||||
two_days_after_tomorrow = (now + datetime.timedelta(days=3)).strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
weekdays = ['星期一', '星期二', '星期三', '星期四', '星期五', '星期六', '星期日']
|
||||
current_weekday = weekdays[now.weekday()]
|
||||
yesterday_weekday = weekdays[(now.weekday() - 1) % 7]
|
||||
tomorrow_weekday = weekdays[(now.weekday() + 1) % 7]
|
||||
day_after_tomorrow_weekday = weekdays[(now.weekday() + 2) % 7]
|
||||
two_days_after_tomorrow_weekday = weekdays[(now.weekday() + 3) % 7]
|
||||
|
||||
result = {
|
||||
"当前时间": f"{current_time} {current_weekday}",
|
||||
"昨天": f"{yesterday} {yesterday_weekday}",
|
||||
"明天": f"{tomorrow} {tomorrow_weekday}",
|
||||
"后天": f"{day_after_tomorrow} {day_after_tomorrow_weekday}",
|
||||
"大后天": f"{two_days_after_tomorrow} {two_days_after_tomorrow_weekday}"
|
||||
}
|
||||
|
||||
return result
|
||||
@@ -0,0 +1,16 @@
|
||||
from typing import Annotated
|
||||
|
||||
from langchain_core.tools import tool
|
||||
|
||||
|
||||
@tool(
|
||||
name_or_callable="tool_multiply",
|
||||
description="一个乘法工具,用于计算两个数字相乘"
|
||||
)
|
||||
def tool_multiply(
|
||||
number1: Annotated[float, "第一个要相乘的数字"],
|
||||
number2: Annotated[float, "第二个要相乘的数字"],
|
||||
) -> float:
|
||||
print("\n invoke tool_multiply", number1, number2, '\n')
|
||||
|
||||
return float(number1) * float(number2)
|
||||
Reference in New Issue
Block a user