feat: 拆分tool文件

This commit is contained in:
martsforever
2025-09-04 21:55:47 +08:00
parent e44c31a54d
commit b123fdb3ec
5 changed files with 131 additions and 97 deletions
+6 -97
View File
@@ -1,117 +1,26 @@
import datetime
import json import json
import time import time
from typing import Union, Annotated, TypedDict from typing import Union
import httpx
from fastapi import FastAPI, Depends from fastapi import FastAPI, Depends
from langchain_core.messages import AIMessage, ToolMessage 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.graph.state import CompiledStateGraph
from langgraph.prebuilt import create_react_agent from langgraph.prebuilt import create_react_agent
from langgraph.types import interrupt, Command from langgraph.types import Command
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from starlette.responses import StreamingResponse from starlette.responses import StreamingResponse
from app.config.env import env
from app.controller.add_user_route import oauth2_scheme from app.controller.add_user_route import oauth2_scheme
from app.model.ConversationModel import ConversationService 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.db_utils import AsyncSessionDep
from app.utils.llm_utils import create_llm from app.utils.llm_utils import create_llm
from app.utils.postgres_checkpointer import PostgresCheckpointerManager, AsyncPostgresSaverDep 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): class ChatMessage(BaseModel):
id: str = Field(..., description="消息id") id: str = Field(..., description="消息id")
type: str = Field(..., description="消息类型") type: str = Field(..., description="消息类型")
+15
View File
@@ -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)
+63
View File
@@ -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},也可以点击这里查看审批进度。"
}
}
]
+31
View File
@@ -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
+16
View File
@@ -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)