feat: init project

This commit is contained in:
martsforever
2025-08-21 22:38:41 +08:00
commit 6ae03985bd
48 changed files with 7221 additions and 0 deletions
@@ -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
+156
View File
@@ -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)
+69
View File
@@ -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
+77
View File
@@ -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)
+52
View File
@@ -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}
+201
View File
@@ -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
+41
View File
@@ -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",
)
+17
View File
@@ -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")
+19
View File
@@ -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
}
+38
View File
@@ -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,
}
+26
View File
@@ -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"}
+27
View File
@@ -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",
)