feat: init project
This commit is contained in:
@@ -0,0 +1,192 @@
|
||||
import json
|
||||
from operator import add
|
||||
from typing import TypedDict, Annotated, List, Literal
|
||||
|
||||
from fastapi import FastAPI
|
||||
from langchain_core.runnables import RunnableConfig
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
from langgraph.constants import START, END
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.graph.state import CompiledStateGraph
|
||||
from langgraph.types import interrupt, Command
|
||||
from sqlalchemy.orm.sync import update
|
||||
|
||||
from app.model.LgApprove import LgApprove, LgApproveService
|
||||
from app.model.LgMessage import LgMessageService
|
||||
from app.utils.db_utils import AsyncSessionDep
|
||||
from app.utils.next_id import next_id
|
||||
from app.utils.postgres_checkpointer import AsyncPostgresSaverDep
|
||||
|
||||
|
||||
def add_langgraph_approve_route(app: FastAPI):
|
||||
# 提交审批单/创建报销单
|
||||
@app.post("/lg_approve/submit")
|
||||
async def lg_approve_submit(
|
||||
row_dict: dict,
|
||||
checkpointer: AsyncPostgresSaverDep,
|
||||
session: AsyncSessionDep,
|
||||
):
|
||||
thread_id = await next_id()
|
||||
graph = create_graph(checkpointer=checkpointer, session=session)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
graph_state = await graph.ainvoke({"input_remarks": row_dict.get('remarks')}, config=config)
|
||||
print(graph_state)
|
||||
return graph_state
|
||||
|
||||
@app.get("/lg_message/feedback/{thread_id}/{approve_flag}")
|
||||
async def lg_message_feedback(
|
||||
thread_id: str,
|
||||
approve_flag: str,
|
||||
checkpointer: AsyncPostgresSaverDep,
|
||||
session: AsyncSessionDep,
|
||||
):
|
||||
graph = create_graph(checkpointer=checkpointer, session=session)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
return await graph.ainvoke(Command(resume=approve_flag), config=config)
|
||||
|
||||
|
||||
def create_graph(
|
||||
checkpointer: AsyncPostgresSaver,
|
||||
session: AsyncSessionDep,
|
||||
) -> CompiledStateGraph:
|
||||
class ApproveSchema(TypedDict):
|
||||
id: str
|
||||
remarks: str
|
||||
status: str
|
||||
result_content: str
|
||||
|
||||
class StateSchema(TypedDict):
|
||||
# 图的入参,需要一个报销单的备注信息,实际业务场景中,入参起码还需要有报销单申请人的id,报销类型,报销金额,发票信息等等;
|
||||
input_remarks: str
|
||||
# 报销单由图中的节点来创建,插入到数据库
|
||||
approve: ApproveSchema
|
||||
# 图执行的结果标识,成功还是失败
|
||||
approve_flag: bool
|
||||
# 每个节点执行的日志信息可以塞到这个字符串数组中
|
||||
log_list: Annotated[List[str], add]
|
||||
# 执行图的时候,如果往消息表中插入了消息数据,也把这个插入的消息记录到状态中
|
||||
lg_message_list: Annotated[List[dict], add]
|
||||
|
||||
builder = StateGraph(StateSchema)
|
||||
|
||||
# 创建报销单(审批单)
|
||||
async def node_create_approve(state: StateSchema):
|
||||
insert_approve_dict = {
|
||||
"status": "pending_approval",
|
||||
"result_content": '待审批......',
|
||||
"remarks": state.get('input_remarks'),
|
||||
}
|
||||
insert_approve_cls = await LgApproveService.item_insert(session=session, row_dict=insert_approve_dict)
|
||||
return {
|
||||
"approve": insert_approve_cls.model_dump(),
|
||||
"log_list": [f"node_create_approve:创建报销单[{insert_approve_cls.id}]"]
|
||||
}
|
||||
|
||||
# 创建消息(提示用户审批)
|
||||
async def node_create_message(state: StateSchema, config: RunnableConfig):
|
||||
thread_id = config.get('configurable').get('thread_id')
|
||||
|
||||
insert_message_dict = {
|
||||
"title": "您有一条报销单待审批。",
|
||||
"content": f"您的下属员工「XXX」提交了一份报销单,报销内容为:{state.get('approve').get('remarks')}",
|
||||
"status": "pending",
|
||||
"render_configs": json.dumps([
|
||||
{
|
||||
"type": "button",
|
||||
"data": {
|
||||
"label": "审批通过",
|
||||
"type": "primary",
|
||||
"submit_url": f"/lg_message/feedback/{thread_id}/Y",
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "button",
|
||||
"data": {
|
||||
"label": "审批拒绝",
|
||||
"submit_url": f"/lg_message/feedback/{thread_id}/N",
|
||||
}
|
||||
},
|
||||
], ensure_ascii=False),
|
||||
}
|
||||
|
||||
insert_message_cls = await LgMessageService.item_insert(session=session, row_dict=insert_message_dict)
|
||||
|
||||
return {
|
||||
"log_list": [f"node_create_message:创建审批消息,待商机主管「XXX」审批"],
|
||||
"lg_message_list": [insert_message_cls.model_dump()],
|
||||
}
|
||||
|
||||
# 触发中断(等待审批恢复重新执行这个节点)
|
||||
async def node_wait_for_approve(state: StateSchema) -> Command[Literal["node_approve_accept", "node_approve_reject"]]:
|
||||
|
||||
approve_flag = interrupt({})
|
||||
|
||||
# 下面是中断恢复的代码
|
||||
|
||||
state_message_dict = state.get('lg_message_list')[-1]
|
||||
update_message_dict = {
|
||||
"id": state_message_dict.get('id'),
|
||||
"status": "proceeded"
|
||||
}
|
||||
update_message_cls = await LgMessageService.item_update(session=session, row_dict=update_message_dict)
|
||||
update_message_dict = update_message_cls.model_dump()
|
||||
|
||||
if approve_flag == 'Y':
|
||||
return Command(
|
||||
goto="node_approve_accept",
|
||||
update={
|
||||
"log_list": [f"node_wait_for_approve:主管审批通过"],
|
||||
"lg_message_list": [update_message_dict],
|
||||
}
|
||||
)
|
||||
else:
|
||||
return Command(
|
||||
goto="node_approve_reject",
|
||||
update={
|
||||
"log_list": [f"node_wait_for_approve:主管审批拒绝"],
|
||||
"lg_message_list": [update_message_dict],
|
||||
}
|
||||
)
|
||||
|
||||
# 审批通过处理节点
|
||||
async def node_approve_accept(state: StateSchema):
|
||||
update_approve_dict = {
|
||||
"id": state.get('approve').get('id'),
|
||||
"status": "accept_approval",
|
||||
"result_content": '审批通过......',
|
||||
}
|
||||
update_approve_cls = await LgApproveService.item_update(session=session, row_dict=update_approve_dict)
|
||||
return {
|
||||
"approve": update_approve_cls.model_dump(),
|
||||
"log_list": [f"node_approve_accept:审批通过,等待财务打款"],
|
||||
}
|
||||
|
||||
# 审批拒绝处理节点
|
||||
async def node_approve_reject(state: StateSchema):
|
||||
update_approve_dict = {
|
||||
"id": state.get('approve').get('id'),
|
||||
"status": "reject_approval",
|
||||
"result_content": '审批拒绝......',
|
||||
}
|
||||
update_approve_cls = await LgApproveService.item_update(session=session, row_dict=update_approve_dict)
|
||||
return {
|
||||
"approve": update_approve_cls.model_dump(),
|
||||
"log_list": [f"node_approve_accept:审批已经被拒绝"],
|
||||
}
|
||||
|
||||
builder.add_node(node_create_approve)
|
||||
builder.add_node(node_create_message)
|
||||
builder.add_node(node_wait_for_approve)
|
||||
builder.add_node(node_approve_accept)
|
||||
builder.add_node(node_approve_reject)
|
||||
|
||||
builder.add_edge(START, 'node_create_approve')
|
||||
builder.add_edge("node_create_approve", 'node_create_message')
|
||||
builder.add_edge("node_create_message", 'node_wait_for_approve')
|
||||
|
||||
builder.add_edge('node_approve_accept', END)
|
||||
builder.add_edge('node_approve_reject', END)
|
||||
|
||||
graph = builder.compile(checkpointer=checkpointer)
|
||||
|
||||
return graph
|
||||
@@ -0,0 +1,156 @@
|
||||
import datetime
|
||||
import json
|
||||
from typing import Union, Annotated
|
||||
|
||||
from fastapi import FastAPI
|
||||
from langchain_core.messages import HumanMessage
|
||||
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 app.utils.llm_utils import create_llm
|
||||
from app.utils.postgres_checkpointer import PostgresCheckpointerManager
|
||||
|
||||
|
||||
@tool(name_or_callable="tool_get_datetime", description="一个用于获取当前时间的工具,没有参数")
|
||||
def tool_get_datetime():
|
||||
return datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
|
||||
@tool(name_or_callable="tool_book_hotel", description="一个用于预定酒店的工具")
|
||||
def tool_book_hotel(
|
||||
hotel_name: Annotated[str, '酒店名称'],
|
||||
room_type: Annotated[str, '房间类型'],
|
||||
check_in_date: Annotated[str, '入住时间'],
|
||||
) -> str:
|
||||
resume_data = interrupt({
|
||||
"title": "请确认酒店预定信息",
|
||||
"form": [
|
||||
{
|
||||
"field": "hotel_name",
|
||||
"type": "input",
|
||||
"label": "酒店名称",
|
||||
"required": True,
|
||||
},
|
||||
{
|
||||
"field": "room_type",
|
||||
"type": "select",
|
||||
"label": "客房类型",
|
||||
"options": [
|
||||
{"label": "标间", "value": "标间"},
|
||||
{"label": "单间", "value": "单间"},
|
||||
{"label": "双人间", "value": "双人间"},
|
||||
],
|
||||
"required": True,
|
||||
},
|
||||
{
|
||||
"field": "check_in_date",
|
||||
"type": "date",
|
||||
"label": "入住时间",
|
||||
}
|
||||
],
|
||||
"formData": {
|
||||
"hotel_name": hotel_name,
|
||||
"room_type": room_type,
|
||||
"check_in_date": check_in_date,
|
||||
}
|
||||
})
|
||||
if resume_data == "N":
|
||||
return f"用户选择取消预定酒店"
|
||||
|
||||
old_book_data = {
|
||||
"hotel_name": hotel_name,
|
||||
"room_type": room_type,
|
||||
"check_in_date": check_in_date,
|
||||
}
|
||||
new_book_data: dict = resume_data
|
||||
|
||||
change_field_list = [(k, v) for k, v in old_book_data.items() if old_book_data.get(k) != new_book_data.get(k)]
|
||||
|
||||
return (f"用户已经更改参数,最新的为{json.dumps(new_book_data, ensure_ascii=False)}" if change_field_list else "") + ",结果为预定成功"
|
||||
|
||||
|
||||
class ChatMessage(BaseModel):
|
||||
id: str = Field(..., description="消息id")
|
||||
type: str = Field(..., description="消息类型")
|
||||
content: str = Field(..., description="消息内容")
|
||||
|
||||
|
||||
# 对话接口参数类型
|
||||
class ChatParam(BaseModel):
|
||||
thread_id: str = Field(..., description="线程id")
|
||||
human_message: ChatMessage = Field(..., description="用户消息")
|
||||
|
||||
|
||||
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(),
|
||||
tools=[tool_book_hotel, tool_get_datetime],
|
||||
checkpointer=await PostgresCheckpointerManager.get_instance(),
|
||||
prompt="""
|
||||
你是一名擅长使用工具的智能助手,你需要根据用户问题来进行回答,请使用中文进行回答。
|
||||
当用户问题需要调用工具时再调用工具,否则按照你的知识来回答问题。
|
||||
某些工具会触发中断让用户来编辑工具执行参数,这些工具会将新的执行参数作为信息返回,你需要回复用户最新的信息
|
||||
"""
|
||||
)
|
||||
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
|
||||
async def get_chat_state(thread_id: str):
|
||||
graph = await ChatAgent.get_agent()
|
||||
graph_state = await graph.aget_state(config={"configurable": {"thread_id": thread_id}})
|
||||
if graph_state.values.get('messages', None) is None:
|
||||
graph_state.values['messages'] = []
|
||||
return {
|
||||
**graph_state.values,
|
||||
"__interrupt__": graph_state.interrupts,
|
||||
}
|
||||
|
||||
|
||||
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/chat_resume/{thread_id}")
|
||||
async def langgraph_chat(body: dict, 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(
|
||||
Command(resume=body.get('resume_data')),
|
||||
config={"configurable": {"thread_id": thread_id}}
|
||||
)
|
||||
return {
|
||||
**graph_state,
|
||||
# 这里不需要加1,因为我们并没有往messages中增加消息
|
||||
"messages": graph_state.get('messages')[len(chat_history_list):],
|
||||
}
|
||||
|
||||
# 查询聊天记录
|
||||
@app.get("/langgraph/chat_state/{thread_id}")
|
||||
async def langgraph_chat(thread_id: str):
|
||||
return await ChatAgent.get_chat_state(thread_id)
|
||||
@@ -0,0 +1,69 @@
|
||||
import random
|
||||
from operator import add
|
||||
from typing import TypedDict, Annotated, List
|
||||
|
||||
from fastapi import FastAPI
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
from langgraph.constants import START, END
|
||||
from langgraph.graph import StateGraph
|
||||
|
||||
from app.utils.postgres_checkpointer import AsyncPostgresSaverDep
|
||||
|
||||
|
||||
def add_langgraph_route(app: FastAPI):
|
||||
@app.get("/langgraph/invoke")
|
||||
async def langgraph_invoke(thread_id: str, checkpointer: AsyncPostgresSaverDep):
|
||||
print("checkpointer", checkpointer)
|
||||
graph = create_graph(checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
graph_state = await graph.ainvoke({"name_list": [f"initial:{thread_id}"]}, config=config)
|
||||
return graph_state
|
||||
|
||||
@app.get("/langgraph/get_state")
|
||||
async def langgraph_get_state(thread_id: str, checkpointer: AsyncPostgresSaverDep):
|
||||
graph = create_graph(checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
graph_state = await graph.aget_state(config)
|
||||
return graph_state.values
|
||||
|
||||
@app.get("/langgraph/get_state_snapshot")
|
||||
async def langgraph_get_state_snapshot(thread_id: str, checkpointer: AsyncPostgresSaverDep):
|
||||
graph = create_graph(checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
graph_state = await graph.aget_state(config)
|
||||
return graph_state
|
||||
|
||||
|
||||
def create_graph(checkpointer: AsyncPostgresSaver):
|
||||
class StateSchema(TypedDict):
|
||||
name_list: Annotated[List[str], add]
|
||||
|
||||
builder = StateGraph(StateSchema)
|
||||
|
||||
def node_1(state):
|
||||
random_int = random.randint(0, 100)
|
||||
print(["🧠节点执行", "node_1", random_int])
|
||||
return {"name_list": [f"node_1:{random_int}"]}
|
||||
|
||||
def node_2(state):
|
||||
random_int = random.randint(100, 200)
|
||||
print(["🧠节点执行", "node_2", random_int])
|
||||
return {"name_list": [f"node_2:{random_int}"]}
|
||||
|
||||
def node_3(state):
|
||||
random_int = random.randint(300, 400)
|
||||
print(["🧠节点执行", "node_3", random_int])
|
||||
return {"name_list": [f"node_3:{random_int}"]}
|
||||
|
||||
builder.add_node(node_1)
|
||||
builder.add_node(node_2)
|
||||
builder.add_node(node_3)
|
||||
|
||||
builder.add_edge(START, 'node_1')
|
||||
builder.add_edge('node_1', 'node_2')
|
||||
builder.add_edge('node_2', 'node_3')
|
||||
builder.add_edge('node_3', END)
|
||||
|
||||
graph = builder.compile(checkpointer=checkpointer)
|
||||
|
||||
return graph
|
||||
@@ -0,0 +1,95 @@
|
||||
import random
|
||||
from contextlib import asynccontextmanager
|
||||
from operator import add
|
||||
from typing import TypedDict, Annotated, List
|
||||
|
||||
from fastapi import FastAPI, Depends
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
from langgraph.constants import START, END
|
||||
from langgraph.graph import StateGraph
|
||||
|
||||
from app.utils.postgres_checkpointer import AsyncPostgresSaverDep
|
||||
|
||||
|
||||
class CustomConnection:
|
||||
def __init__(self):
|
||||
print("开始自定义连接")
|
||||
|
||||
def close(self):
|
||||
print("关闭自定义连接")
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def open_custom_connection():
|
||||
ctx_connection = CustomConnection()
|
||||
yield ctx_connection
|
||||
ctx_connection.close()
|
||||
|
||||
|
||||
async def get_async_ctx_conn():
|
||||
async with open_custom_connection() as conn:
|
||||
yield conn
|
||||
|
||||
|
||||
CustomConnDep = Annotated[CustomConnection, Depends(get_async_ctx_conn)]
|
||||
|
||||
|
||||
def add_langgraph_route(app: FastAPI):
|
||||
@app.get("/langgraph/invoke")
|
||||
async def langgraph_invoke(thread_id: str, checkpointer: AsyncPostgresSaverDep, cst_conn: CustomConnDep):
|
||||
print("start", cst_conn)
|
||||
print("checkpointer", checkpointer)
|
||||
graph = create_graph(checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
graph_state = await graph.ainvoke({"name_list": [f"initial:{thread_id}"]}, config=config)
|
||||
print("end")
|
||||
return graph_state
|
||||
|
||||
@app.get("/langgraph/get_state")
|
||||
async def langgraph_get_state(thread_id: str, checkpointer: AsyncPostgresSaverDep):
|
||||
graph = create_graph(checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
graph_state = await graph.aget_state(config)
|
||||
return graph_state.values
|
||||
|
||||
@app.get("/langgraph/get_state_snapshot")
|
||||
async def langgraph_get_state_snapshot(thread_id: str, checkpointer: AsyncPostgresSaverDep):
|
||||
graph = create_graph(checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
graph_state = await graph.aget_state(config)
|
||||
return graph_state
|
||||
|
||||
|
||||
def create_graph(checkpointer: AsyncPostgresSaver):
|
||||
class StateSchema(TypedDict):
|
||||
name_list: Annotated[List[str], add]
|
||||
|
||||
builder = StateGraph(StateSchema)
|
||||
|
||||
def node_1(state):
|
||||
random_int = random.randint(0, 100)
|
||||
print(["🧠节点执行", "node_1", random_int])
|
||||
return {"name_list": [f"node_1:{random_int}"]}
|
||||
|
||||
def node_2(state):
|
||||
random_int = random.randint(100, 200)
|
||||
print(["🧠节点执行", "node_2", random_int])
|
||||
return {"name_list": [f"node_2:{random_int}"]}
|
||||
|
||||
def node_3(state):
|
||||
random_int = random.randint(300, 400)
|
||||
print(["🧠节点执行", "node_3", random_int])
|
||||
return {"name_list": [f"node_3:{random_int}"]}
|
||||
|
||||
builder.add_node(node_1)
|
||||
builder.add_node(node_2)
|
||||
builder.add_node(node_3)
|
||||
|
||||
builder.add_edge(START, 'node_1')
|
||||
builder.add_edge('node_1', 'node_2')
|
||||
builder.add_edge('node_2', 'node_3')
|
||||
builder.add_edge('node_3', END)
|
||||
|
||||
graph = builder.compile(checkpointer=checkpointer)
|
||||
|
||||
return graph
|
||||
@@ -0,0 +1,77 @@
|
||||
import random
|
||||
from operator import add
|
||||
from typing import Annotated, List, TypedDict
|
||||
|
||||
from fastapi import FastAPI
|
||||
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
|
||||
from langgraph.constants import START, END
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.types import interrupt, Command
|
||||
|
||||
from app.utils.postgres_checkpointer import AsyncPostgresSaverDep
|
||||
|
||||
|
||||
def add_lg_approve_route(app: FastAPI):
|
||||
@app.get("/lg/approve/submit")
|
||||
async def lg_approve_submit(thread_id: str, checkpointer: AsyncPostgresSaverDep):
|
||||
graph = create_graph(checkpointer=checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
return await graph.ainvoke({"name_list": [f"initial:{thread_id}"]}, config=config)
|
||||
|
||||
@app.get("/lg/approve/state")
|
||||
async def lg_approve_get_state(thread_id: str, checkpointer: AsyncPostgresSaverDep):
|
||||
graph = create_graph(checkpointer=checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
state_snapshot = await graph.aget_state(config)
|
||||
graph_state = state_snapshot.values
|
||||
return {
|
||||
**graph_state,
|
||||
"__interrupt__": state_snapshot.interrupts,
|
||||
}
|
||||
|
||||
@app.get("/lg/approve/resume")
|
||||
async def lg_approve_resume(
|
||||
thread_id: str,
|
||||
is_approve: str,
|
||||
checkpointer: AsyncPostgresSaverDep
|
||||
):
|
||||
graph = create_graph(checkpointer=checkpointer)
|
||||
config = {"configurable": {"thread_id": thread_id}}
|
||||
return await graph.ainvoke(Command(resume=is_approve), config=config)
|
||||
|
||||
|
||||
def create_graph(checkpointer: AsyncPostgresSaver):
|
||||
class StateSchema(TypedDict):
|
||||
name_list: Annotated[List[str], add]
|
||||
|
||||
builder = StateGraph(StateSchema)
|
||||
|
||||
def node_1(state: StateSchema):
|
||||
random_int = random.randint(0, 100)
|
||||
print(["🧠节点执行", "node_1", random_int])
|
||||
return {"name_list": [f"node_1:{random_int}"]}
|
||||
|
||||
def node_2(state: StateSchema):
|
||||
random_int = random.randint(100, 200)
|
||||
print(["🧠节点执行", "node_2", random_int])
|
||||
is_approved = interrupt({
|
||||
"message": "需要主管审批"
|
||||
})
|
||||
result = "✅通过" if is_approved == 'Y' else "❌拒绝"
|
||||
return {"name_list": [f"node_2:{random_int}," + result]}
|
||||
|
||||
def node_3(state: StateSchema):
|
||||
random_int = random.randint(200, 300)
|
||||
print(["🧠节点执行", "node_3", random_int])
|
||||
return {"name_list": [f"node_3:{random_int}"]}
|
||||
|
||||
builder.add_node(node_1)
|
||||
builder.add_node(node_2)
|
||||
builder.add_node(node_3)
|
||||
|
||||
builder.add_edge(START, 'node_1')
|
||||
builder.add_edge("node_1", 'node_2')
|
||||
builder.add_edge("node_2", 'node_3')
|
||||
builder.add_edge("node_3", END)
|
||||
|
||||
return builder.compile(checkpointer=checkpointer)
|
||||
@@ -0,0 +1,52 @@
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from sqlmodel import select
|
||||
|
||||
from app.model.LlmUser import LlmUser
|
||||
from app.utils.db_utils import AsyncSessionDep
|
||||
from app.utils.next_id import next_id
|
||||
|
||||
|
||||
def add_sqlmodel_route(app: FastAPI):
|
||||
@app.post("/llm_user/insert")
|
||||
async def llm_user_insert(user: LlmUser, session: AsyncSessionDep):
|
||||
if user.id is None:
|
||||
user.id = await next_id()
|
||||
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
return {"result": user}
|
||||
|
||||
@app.post("/llm_user/update")
|
||||
async def llm_user_insert(user_dict: dict, session: AsyncSessionDep):
|
||||
|
||||
if user_dict.get("id") is None:
|
||||
raise HTTPException(status_code=500, detail="Update row missing id")
|
||||
|
||||
update_user = (await session.exec(select(LlmUser).where(LlmUser.id == user_dict["id"]))).first()
|
||||
if update_user is None:
|
||||
raise HTTPException(status_code=500, detail="Update row not found")
|
||||
|
||||
for key, value in user_dict.items():
|
||||
setattr(update_user, key, value)
|
||||
|
||||
session.add(update_user)
|
||||
await session.commit()
|
||||
await session.refresh(update_user)
|
||||
return {"result": update_user}
|
||||
|
||||
@app.post("/llm_user/delete")
|
||||
async def llm_user_delete(user_dict: dict, session: AsyncSessionDep):
|
||||
|
||||
if user_dict.get("id") is None:
|
||||
raise HTTPException(status_code=500, detail="Update row missing id")
|
||||
|
||||
delete_user: LlmUser = (await session.exec(select(LlmUser).where(LlmUser.id == user_dict["id"]))).first()
|
||||
|
||||
if delete_user is None:
|
||||
raise HTTPException(status_code=500, detail="Delete row not found")
|
||||
|
||||
await session.delete(delete_user)
|
||||
await session.commit()
|
||||
|
||||
return {"result": True}
|
||||
@@ -0,0 +1,201 @@
|
||||
from datetime import timedelta
|
||||
from enum import Enum
|
||||
|
||||
from fastapi import FastAPI, Depends, HTTPException
|
||||
from fastapi.security import OAuth2PasswordRequestForm, OAuth2PasswordBearer
|
||||
from jwt import InvalidTokenError
|
||||
from pydantic import BaseModel
|
||||
from sqlmodel import select, Field
|
||||
from starlette import status
|
||||
|
||||
from app.config.env import env
|
||||
from app.model.BasicModel import BasicModel
|
||||
from app.utils.CrpyUtils import CryptUtils
|
||||
from app.utils.db_utils import AsyncSessionDep
|
||||
from app.utils.next_id import next_id
|
||||
|
||||
|
||||
class UserValidate(str, Enum):
|
||||
Y = 'Y'
|
||||
N = 'N'
|
||||
|
||||
|
||||
# 公共的,也是最后返回给前端的一个用户信息数据类型
|
||||
class PublicUser(BasicModel):
|
||||
username: str = Field(..., description="用户名")
|
||||
email: str = Field(..., description="邮箱")
|
||||
full_name: str = Field(..., description="用户全名")
|
||||
valid: UserValidate = Field(default=UserValidate.N, description="用户账号是否已经激活")
|
||||
|
||||
|
||||
# 注册的时候,客户端传入的用户信息,需要包含这个明文密码字段
|
||||
class RegistryUser(PublicUser):
|
||||
password: str
|
||||
|
||||
|
||||
# 对pl_user表进行增删改查时的这个model类
|
||||
class UserModel(PublicUser, table=True):
|
||||
__tablename__ = "pl_user"
|
||||
hash_password: str
|
||||
|
||||
|
||||
class Token(BaseModel):
|
||||
token: str
|
||||
token_type: str
|
||||
|
||||
|
||||
def add_user_route(app: FastAPI):
|
||||
@app.post("/registry")
|
||||
async def _registry(registry_user: RegistryUser, session: AsyncSessionDep):
|
||||
|
||||
# /*---------------------------------------检查用户名是否已经注册-------------------------------------------*/
|
||||
query = select(UserModel).where(UserModel.username == registry_user.username)
|
||||
result = await session.execute(query)
|
||||
item_cls = result.scalars().first()
|
||||
|
||||
if item_cls:
|
||||
return {"result": None, "error": f"用户名:{registry_user.username} 已经存在"}
|
||||
|
||||
# /*---------------------------------------检查邮箱是否已经注册-------------------------------------------*/
|
||||
|
||||
query = select(UserModel).where(UserModel.email == registry_user.email)
|
||||
result = await session.execute(query)
|
||||
item_cls = result.scalars().first()
|
||||
|
||||
if item_cls:
|
||||
return {"result": None, "error": f"邮箱:{registry_user.email} 已经注册"}
|
||||
|
||||
# /*---------------------------------------开始注册流程-------------------------------------------*/
|
||||
|
||||
hash_password = CryptUtils.get_password_hash(registry_user.password)
|
||||
user = UserModel(
|
||||
username=registry_user.username,
|
||||
email=registry_user.email,
|
||||
full_name=registry_user.full_name,
|
||||
hash_password=hash_password,
|
||||
valid=UserValidate.N,
|
||||
)
|
||||
user.id = await next_id()
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
|
||||
public_user = PublicUser(**user.model_dump())
|
||||
|
||||
active_user_token = CryptUtils.create_access_token(public_user.username, expires_delta=timedelta(days=365 * 3))
|
||||
active_url = f"{env.server_domain}:{env.server_port}/verify?token={active_user_token}"
|
||||
|
||||
return {
|
||||
"result": public_user,
|
||||
"active_url": active_url
|
||||
}
|
||||
|
||||
@app.get("/verify")
|
||||
async def _verify(token: str, session: AsyncSessionDep):
|
||||
username = CryptUtils.get_username_from_token(token)
|
||||
|
||||
if not username:
|
||||
return {"result": None, "error": "token无效或者已经过期"}
|
||||
|
||||
query = select(UserModel).where(UserModel.username == username)
|
||||
result = await session.execute(query)
|
||||
item_cls: UserModel | None = result.scalars().first()
|
||||
|
||||
if not item_cls:
|
||||
return {"result": None, "error": f"用户 {username} 不存在"}
|
||||
|
||||
item_cls.valid = UserValidate.Y
|
||||
session.add(item_cls)
|
||||
await session.commit()
|
||||
await session.refresh(item_cls)
|
||||
|
||||
public_user = PublicUser(**item_cls.model_dump())
|
||||
|
||||
return {
|
||||
"result": public_user,
|
||||
"message": f"用户 {username} 激活成功"
|
||||
}
|
||||
|
||||
@app.post("/login")
|
||||
@app.post("/token")
|
||||
async def _token(session: AsyncSessionDep, form_data: OAuth2PasswordRequestForm = Depends()):
|
||||
print("login", form_data)
|
||||
user = await authenticate_user(session, form_data.username, form_data.password)
|
||||
|
||||
if not user:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="用户名或者密码不正确",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
token = Token(
|
||||
token=CryptUtils.create_access_token(user.username),
|
||||
token_type="Bearer",
|
||||
)
|
||||
|
||||
return {
|
||||
"result": user,
|
||||
"token": token,
|
||||
}
|
||||
|
||||
@app.get("/users/me")
|
||||
async def _me(current_user: PublicUser = Depends(get_current_user)):
|
||||
return current_user
|
||||
|
||||
@app.post("/order")
|
||||
async def _query_order(product_name: str, current_user: PublicUser = Depends(get_current_user)):
|
||||
return [product_name]
|
||||
|
||||
|
||||
async def authenticate_user(session: AsyncSessionDep, username: str, password: str):
|
||||
query = select(UserModel).where(UserModel.username == username)
|
||||
result = await session.execute(query)
|
||||
item_cls: UserModel | None = result.scalars().first()
|
||||
if not item_cls:
|
||||
return None
|
||||
|
||||
if item_cls.valid != UserValidate.Y:
|
||||
return None
|
||||
|
||||
if not CryptUtils.verify_password(password, item_cls.hash_password):
|
||||
return None
|
||||
|
||||
public_user = PublicUser(**item_cls.model_dump())
|
||||
|
||||
return public_user
|
||||
|
||||
|
||||
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="token")
|
||||
|
||||
unauthorized_exception = HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail="The token is invalid or had expired",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
|
||||
async def get_current_user(session: AsyncSessionDep, token: str = Depends(oauth2_scheme)):
|
||||
try:
|
||||
username = CryptUtils.get_username_from_token(token)
|
||||
if not username:
|
||||
raise unauthorized_exception
|
||||
except InvalidTokenError:
|
||||
raise unauthorized_exception
|
||||
|
||||
user_model = await get_user_by_username(username, session)
|
||||
if not user_model:
|
||||
raise unauthorized_exception
|
||||
|
||||
return PublicUser(**user_model.model_dump())
|
||||
|
||||
|
||||
async def get_user_by_username(username: str, session: AsyncSessionDep):
|
||||
query = (
|
||||
select(UserModel)
|
||||
.where(UserModel.username == username)
|
||||
.where(UserModel.valid == UserValidate.Y)
|
||||
)
|
||||
result = await session.execute(query)
|
||||
item_cls: UserModel | None = result.scalars().first()
|
||||
return item_cls
|
||||
@@ -0,0 +1,41 @@
|
||||
from fastapi import FastAPI
|
||||
from langchain_core.output_parsers import StrOutputParser
|
||||
from langchain_core.prompts import ChatPromptTemplate, MessagesPlaceholder
|
||||
from langserve import add_routes
|
||||
from pydantic import Field
|
||||
|
||||
from app.utils.ModelInputSchema import ModelInputSchema
|
||||
from app.utils.llm_utils import create_llm
|
||||
|
||||
|
||||
class ModelInputSchema2(ModelInputSchema):
|
||||
language: str = Field(..., description="要求模型回答使用的语言")
|
||||
|
||||
|
||||
_doubao_chain2 = (
|
||||
ChatPromptTemplate.from_messages([
|
||||
('system', "你需要使用语言“{language}”来回答用户的问题, 不论如何,你必须使用“{language}”来回答用户"),
|
||||
MessagesPlaceholder(variable_name="messages")
|
||||
]) |
|
||||
create_llm() |
|
||||
StrOutputParser()
|
||||
)
|
||||
|
||||
|
||||
def add_custom_chat_playground_route(app: FastAPI):
|
||||
add_routes(
|
||||
app=app,
|
||||
runnable=_doubao_chain2,
|
||||
input_type=ModelInputSchema2,
|
||||
path="/doubao2"
|
||||
)
|
||||
|
||||
add_routes(
|
||||
app=app,
|
||||
runnable={
|
||||
"messages": lambda x: x['messages'],
|
||||
"language": lambda x: "英语"
|
||||
} | _doubao_chain2,
|
||||
input_type=ModelInputSchema,
|
||||
path="/doubao2_playgroud",
|
||||
)
|
||||
@@ -0,0 +1,17 @@
|
||||
import asyncio
|
||||
import json
|
||||
|
||||
from starlette.responses import StreamingResponse
|
||||
|
||||
|
||||
def add_custom_stream_api_route(app):
|
||||
@app.get("/my_stream")
|
||||
async def custom_stream(start: int, end: int):
|
||||
async def generate_numbers():
|
||||
current = start
|
||||
while current <= end:
|
||||
yield json.dumps({"number": current}) + "\n"
|
||||
await asyncio.sleep(0.2)
|
||||
current += 1
|
||||
|
||||
return StreamingResponse(generate_numbers(), media_type="application/x-ndjson")
|
||||
@@ -0,0 +1,19 @@
|
||||
from fastapi import FastAPI
|
||||
from sqlalchemy.sql.expression import text
|
||||
|
||||
from app.utils.db_utils import AsyncSessionDep
|
||||
|
||||
|
||||
def add_test_connection_route(app: FastAPI):
|
||||
@app.get("/query_llm_user_list")
|
||||
async def query_llm_user_list(session: AsyncSessionDep):
|
||||
result = await session.execute(text("select * from llm_user"))
|
||||
return [dict(row._mapping) for row in result]
|
||||
|
||||
@app.get("/query_llm_user")
|
||||
async def query_llm_user(username: str, session: AsyncSessionDep):
|
||||
result = await session.execute(text("select * from llm_user where username = :username"), {"username": username})
|
||||
list = [dict(row._mapping) for row in result]
|
||||
return {
|
||||
"result": list[0] if list else None
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import FastAPI
|
||||
from sqlmodel import select, or_, and_
|
||||
|
||||
from app.model.LlmProduct import LlmProduct
|
||||
from app.model.LlmUser import LlmUser
|
||||
from app.utils.db_utils import AsyncSessionDep
|
||||
|
||||
|
||||
def add_test_sqlmodel_route(app: FastAPI):
|
||||
@app.get("/llm_user_list")
|
||||
async def llm_user_list(session: AsyncSessionDep):
|
||||
query = select(LlmProduct)
|
||||
query = query.where(
|
||||
or_(
|
||||
and_(
|
||||
LlmProduct.name.not_in(['手机', '电脑', '相机']),
|
||||
LlmProduct.price > 1000,
|
||||
),
|
||||
LlmProduct.price == 300,
|
||||
)
|
||||
)
|
||||
result = await session.execute(query)
|
||||
|
||||
print(type(result), result)
|
||||
return {
|
||||
# "first": result.scalars().first(),
|
||||
"all": result.scalars().all(),
|
||||
}
|
||||
|
||||
@app.post("/llm_user")
|
||||
async def llm_user(product_dict: dict, session: AsyncSessionDep):
|
||||
_product_cls = LlmProduct.model_validate(product_dict)
|
||||
return {
|
||||
"product": product_dict,
|
||||
"_product": _product_cls,
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
import asyncio
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
|
||||
def add_test_sync_route(app: FastAPI):
|
||||
@app.get("/test")
|
||||
async def test():
|
||||
print(f"Process {os.getpid()} handling /test")
|
||||
return {"message": "Hello World"}
|
||||
|
||||
@app.get("/sync_delay")
|
||||
async def sync_delay(delay: int = 1):
|
||||
"""同步延迟delay秒"""
|
||||
print(f"Process {os.getpid()} handling /sync_delay")
|
||||
time.sleep(delay)
|
||||
return {"hello": "world"}
|
||||
|
||||
@app.get("/async_delay")
|
||||
async def async_delay(delay: int = 1):
|
||||
"""异步延迟delay秒"""
|
||||
print(f"Process {os.getpid()} handling /async_delay")
|
||||
await asyncio.sleep(delay)
|
||||
return {"hello": "world"}
|
||||
@@ -0,0 +1,27 @@
|
||||
from fastapi import FastAPI
|
||||
from langchain_core.output_parsers import StrOutputParser
|
||||
from langchain_core.prompts import ChatPromptTemplate
|
||||
from langserve import add_routes
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from app.utils.llm_utils import create_llm
|
||||
|
||||
|
||||
class TranslateChainInputSchema(BaseModel):
|
||||
language: str = Field(..., description="要翻译的目标语言")
|
||||
input: str = Field(..., description="要翻译的内容")
|
||||
|
||||
|
||||
translate_chain = ChatPromptTemplate.from_messages([
|
||||
('system', '你需要把用户的内容翻译为:{language}'),
|
||||
('user', "{input}")
|
||||
]) | create_llm() | StrOutputParser()
|
||||
|
||||
|
||||
def add_translate_route(app: FastAPI):
|
||||
add_routes(
|
||||
app=app,
|
||||
runnable=translate_chain,
|
||||
input_type=TranslateChainInputSchema,
|
||||
path="/translate",
|
||||
)
|
||||
Reference in New Issue
Block a user