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
+23
View File
@@ -0,0 +1,23 @@
# http://editorconfig.org
root = true
[*]
indent_style = space
indent_size = 2
charset = utf-8
trim_trailing_whitespace = true
insert_final_newline = true
max_line_length = off
ij_javascript_use_double_quotes = true
ij_javascript_force_semicolon_style = true
ij_javascript_spaces_within_object_type_braces = true
ij_javascript_spaces_within_object_literal_braces = true
ij_typescript_use_double_quotes = true
ij_typescript_force_semicolon_style = true
ij_typescript_spaces_within_object_type_braces = true
ij_typescript_spaces_within_object_literal_braces = true
[*.md]
trim_trailing_whitespace = false
+25
View File
@@ -0,0 +1,25 @@
DB_HOST=xxx.xxx.xxx.xxx # 数据库连接ip地址
DB_PORT=xxx # 数据库连接端口
DB_USERNAME=xxx # 数据库连接用户名
DB_PASSWORD=xxx # 数据库连接密码
DB_DATABASE=xxx # 数据库连接的数据库名
PG_DB_HOST=xxx.xxx.xxx.xxx # postgres数据库连接ip地址
PG_DB_PORT=xxx # postgres数据库连接端口
PG_DB_USERNAME=xxx # postgres数据库连接用户名
PG_DB_PASSWORD=xxx # postgres数据库连接密码
PG_DB_DATABASE=xxx # postgres数据库连接的数据库名
LLM_KEY_LOCAL=123
LLM_KEY_HUOSHAN=a0311f2a-ba85-4428-b158-xxxxxxxxxxxx # 火山引擎模型服务平台key
LLM_KEY_BAILIAN=sk-51d13ba8ea044d128c66dxxxxxxxxxxxx # 阿里云百炼模型服务平台key
LLM_KEY_DEEPSEEK=sk-a89d0ff9421a43fca5f0xxxxxxxxxxxx # Deepseek模型服务平台key
SERVER_PORT = 7002 # 服务启动端口
SERVER_DOMAIN = http://127.0.0.1 # 后端服务部署的域名
JWT_SECRET_KEY = a2d41c19c766490a2bea4be9902ea9f60e44ca3d9ef82e8b803f77a4f9deb247 # JWT秘钥
JWT_ALGORITHM = HS256 # JWT加密算法
JWT_ACCESS_TOKEN_EXPIRE_MINUTES = 30 # JWT访问令牌默认有效时间(分钟
JWT_GLOBAL_ENABLE = false # 是否开启全局JWT认证
JWT_WHITE_LIST = ["/token", "/login", "/registry", "/verify", "/async_delay"] # JWT认证白名单接口,不需要认证的接口
+3
View File
@@ -0,0 +1,3 @@
__pycache__
.idea
.env
+21
View File
@@ -0,0 +1,21 @@
FROM python:3.11-slim
RUN pip install poetry==1.6.1
RUN poetry config virtualenvs.create false
WORKDIR /code
COPY ./pyproject.toml ./README.md ./poetry.lock* ./
COPY ./package[s] ./packages
RUN poetry install --no-interaction --no-ansi --no-root
COPY ./app ./app
RUN poetry install --no-interaction --no-ansi
EXPOSE 8080
CMD exec uvicorn app.server:app --host 0.0.0.0 --port 8080
+79
View File
@@ -0,0 +1,79 @@
# ai-langserve
## Installation
Install the LangChain CLI if you haven't yet
```bash
pip install -U langchain-cli
```
## Adding packages
```bash
# adding packages from
# https://github.com/langchain-ai/langchain/tree/master/templates
langchain app add $PROJECT_NAME
# adding custom GitHub repo packages
langchain app add --repo $OWNER/$REPO
# or with whole git string (supports other git providers):
# langchain app add git+https://github.com/hwchase17/chain-of-verification
# with a custom api mount point (defaults to `/{package_name}`)
langchain app add $PROJECT_NAME --api_path=/my/custom/path/rag
```
Note: you remove packages by their api path
```bash
langchain app remove my/custom/path/rag
```
## Setup LangSmith (Optional)
LangSmith will help us trace, monitor and debug LangChain applications.
You can sign up for LangSmith [here](https://smith.langchain.com/).
If you don't have access, you can skip this section
```shell
export LANGSMITH_TRACING=true
export LANGSMITH_API_KEY=<your-api-key>
export LANGSMITH_PROJECT=<your-project> # if not specified, defaults to "default"
```
## Launch LangServe
```bash
langchain serve
```
## Running in Docker
This project folder includes a Dockerfile that allows you to easily build and host your LangServe app.
### Building the Image
To build the image, you simply:
```shell
docker build . -t my-langserve-app
```
If you tag your image with something other than `my-langserve-app`,
note it for use in the next step.
### Running the Image Locally
To run the image, you'll need to include any environment variables
necessary for your application.
In the below example, we inject the `OPENAI_API_KEY` environment
variable with the value set in my local environment
(`$OPENAI_API_KEY`)
We also expose port 8080 with the `-p 8080:8080` option.
```shell
docker run -e OPENAI_API_KEY=$OPENAI_API_KEY -p 8080:8080 my-langserve-app
```
View File
+64
View File
@@ -0,0 +1,64 @@
from app.config.env import env
ai_configs = {
# /*---------------------------------------local-------------------------------------------*/
"local": {
"model": "deepseek-r1-distill-qwen-7b",
"url": "http://127.0.0.1:1234/v1/chat/completions",
"key": env.llm_key_local
},
"local_glm": {
"model": "glm-4-9b-0414",
"url": "http://127.0.0.1:1234/v1/chat/completions",
"key": env.llm_key_local
},
"local_embedding": {
"model": "text-embedding-text2vec-large-chinese",
"url": "http://127.0.0.1:1234/v1/embeddings",
"key": env.llm_key_local
},
# /*---------------------------------------deepseek-------------------------------------------*/
'deepseek-v3': {
'model': 'deepseek-chat',
'url': 'https://api.deepseek.com/chat/completions',
'key': env.llm_key_deepseek,
},
'deepseek-r1': {
'model': 'deepseek-reasoner',
'url': 'https://api.deepseek.com/chat/completions',
'key': env.llm_key_deepseek,
},
# /*---------------------------------------huoshan-------------------------------------------*/
'huoshan-deepseek-v3': {
'model': 'deepseek-v3-250324',
'url': 'https://ark.cn-beijing.volces.com/api/v3/chat/completions',
'key': env.llm_key_huoshan,
},
'huoshan-deepseek-r1': {
'model': 'deepseek-r1-distill-qwen-7b-250120',
'url': 'https://ark.cn-beijing.volces.com/api/v3/chat/completions',
'key': env.llm_key_huoshan,
},
'huoshan-doubao': {
'model': 'doubao-1-5-lite-32k-250115',
'url': 'https://ark.cn-beijing.volces.com/api/v3/chat/completions',
'key': env.llm_key_huoshan,
},
'huoshan-doubao-seed': {
'model': 'doubao-seed-1-6-250615',
'url': 'https://ark.cn-beijing.volces.com/api/v3/chat/completions',
'key': env.llm_key_huoshan,
},
'huoshan-embedding-240715': {
"model": 'doubao-embedding-text-240715',
"url": 'https://ark.cn-beijing.volces.com/api/v3/embeddings',
'key': env.llm_key_huoshan,
},
# /*---------------------------------------bailian-------------------------------------------*/
'bailian-qwen-turbo': {
'model': 'qwen-turbo',
'url': 'https://dashscope.aliyuncs.com/compatible-mode/v1/chat/completions',
'key': env.llm_key_bailian,
},
}
+41
View File
@@ -0,0 +1,41 @@
from typing import List
from pydantic_settings import BaseSettings
from pydantic import Field
from dotenv import load_dotenv
class Settings(BaseSettings):
db_host: str = Field(..., env="DB_HOST")
db_port: str = Field(..., env="DB_PORT")
db_username: str = Field(..., env="DB_USERNAME")
db_password: str = Field(..., env="DB_PASSWORD")
db_database: str = Field(..., env="DB_DATABASE")
pg_db_host: str = Field(..., env="PG_DB_HOST")
pg_db_port: str = Field(..., env="PG_DB_PORT")
pg_db_username: str = Field(..., env="PG_DB_USERNAME")
pg_db_password: str = Field(..., env="PG_DB_PASSWORD")
pg_db_database: str = Field(..., env="PG_DB_DATABASE")
llm_key_local: str = Field(..., env="LLM_KEY_LOCAL")
llm_key_huoshan: str = Field(..., env="LLM_KEY_HUOSHAN")
llm_key_bailian: str = Field(..., env="LLM_KEY_BAILIAN")
llm_key_deepseek: str = Field(..., env="LLM_KEY_DEEPSEEK")
server_port: str = Field(..., env="SERVER_PORT")
server_domain: str = Field(..., env="SERVER_DOMAIN")
jwt_secret_key: str = Field(..., env="JWT_SECRET_KEY")
jwt_algorithm: str = Field(..., env="JWT_ALGORITHM")
jwt_access_token_expire_minutes: int = Field(..., env="JWT_ACCESS_TOKEN_EXPIRE_MINUTES")
jwt_global_enable: bool = Field(..., env="JWT_GLOBAL_ENABLE")
jwt_white_list: List[str] = Field(..., env="JWT_WHITE_LIST")
class Config:
env_file = ".env"
env_file_encoding = "utf-8"
load_dotenv(".env") # 先加载 .env 文件
env = Settings()
@@ -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",
)
+68
View File
@@ -0,0 +1,68 @@
from contextlib import asynccontextmanager
from fastapi import FastAPI
from fastapi.openapi.docs import get_swagger_ui_html, get_swagger_ui_oauth2_redirect_html, get_redoc_html
from starlette.middleware.cors import CORSMiddleware
from starlette.responses import RedirectResponse
from starlette.staticfiles import StaticFiles
from app.utils.db_utils import check_database_connection
from app.utils.postgres_checkpointer import check_postgres_connection, close_postgres_connection
def create_app():
@asynccontextmanager
async def lifespan(app: FastAPI):
print("lifespan:启动阶段")
async_engine = await check_database_connection()
await check_postgres_connection()
yield
print("lifespan:销毁阶段")
await async_engine.dispose()
await close_postgres_connection()
app = FastAPI(
docs_url=None, # 禁用默认 Swagger
redoc_url=None, # 禁用默认 ReDoc
lifespan=lifespan,
)
app.mount("/static", StaticFiles(directory="static"), name="static")
# 自定义 Swagger 页面(使用本地资源)
@app.get("/docs", include_in_schema=False)
async def custom_swagger_ui():
return get_swagger_ui_html(
openapi_url=app.openapi_url,
title=app.title + " - Swagger UI",
oauth2_redirect_url=app.swagger_ui_oauth2_redirect_url,
swagger_js_url="/static/swagger-ui-bundle.min.js",
swagger_css_url="/static/swagger-ui.min.css",
)
@app.get(app.swagger_ui_oauth2_redirect_url, include_in_schema=False)
async def swagger_ui_redirect():
return get_swagger_ui_oauth2_redirect_html()
@app.get("/redoc", include_in_schema=False)
async def redoc_html():
return get_redoc_html(
openapi_url=app.openapi_url,
title=app.title + " - ReDoc",
redoc_js_url="/static/redoc.standalone.js",
)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
expose_headers=["*"],
)
@app.get("/")
async def redirect_root_to_docs():
return RedirectResponse("/docs")
return app
+69
View File
@@ -0,0 +1,69 @@
import time
from fastapi import FastAPI, HTTPException
from passlib.exc import InvalidTokenError
from starlette import status
from starlette.requests import Request
from starlette.responses import JSONResponse
from app.config.env import env
from app.controller.add_user_route import unauthorized_exception, get_current_user
from app.utils.db_utils import async_session
def add_app_middlewares(app: FastAPI):
@app.middleware("http")
async def add_process_time_header(request: Request, call_next):
print("time start")
start_time = time.time()
response = await call_next(request)
process_time = f"time:{time.time() - start_time}s"
response.headers['x-Process-Time'] = process_time
print("time end")
return response
# @app.middleware("http")
# async def middleware2(request: Request, call_next):
# print("middleware2 start")
# response = await call_next(request)
# print("middleware2 end")
# return response
@app.middleware("http")
async def add_oauth_middleware(request: Request, call_next):
if not env.jwt_global_enable:
# 没有开启全局的接口认证功能
return await call_next(request)
# 判断接口是否为认证白名单中的接口
if request.url.path in env.jwt_white_list:
return await call_next(request)
token: str | None = None
oauth_header = request.headers.get("Authorization")
if oauth_header and oauth_header.startswith("Bearer "):
token = oauth_header.split(" ")[1].strip()
if not token:
raise unauthorized_exception
async with async_session() as session:
try:
public_user = await get_current_user(session, token)
request.state.user = public_user
request.state.token = token
except InvalidTokenError:
raise unauthorized_exception
response = await call_next(request)
return response
@app.middleware("http")
async def catch_authorized(request: Request, call_next):
try:
response = await call_next(request)
except HTTPException as e:
if e.status_code == status.HTTP_401_UNAUTHORIZED:
return JSONResponse(content=e.detail, status_code=status.HTTP_401_UNAUTHORIZED)
else:
raise e
return response
+62
View File
@@ -0,0 +1,62 @@
from datetime import datetime, timezone, timedelta, date
from pydantic import model_validator
from sqlmodel import SQLModel, Field
# 定义北京时区(UTC+8)
beijing_timezone = timezone(timedelta(hours=8))
# 定义获取当前北京时区时间的匿名函数,用于默认值生成
current_datetime = lambda: datetime.now(beijing_timezone)
# 定义基础模型类,所有其他模型类的父类,包含通用字段和配置
class BasicModel(SQLModel):
# 唯一标识字段,主键,默认为None(通常由系统生成),描述为“唯一标识,编号”
id: str = Field(default=None, primary_key=True, description="唯一标识,编号")
# 创建时间字段,默认值为当前北京时区时间,描述为“创建时间”
created_at: datetime = Field(default_factory=current_datetime, description="创建时间")
# 更新时间字段,默认值为当前北京时区时间,描述为“更新时间”
updated_at: datetime = Field(default_factory=current_datetime, description="更新时间")
# 创建人ID字段,默认为None,描述为“创建人id”
created_by: str | None = Field(default=None, description="创建人id")
# 更新人ID字段,默认为None,描述为“更新人id”
updated_by: str | None = Field(default=None, description="更新人id")
# 模型配置类,用于设置JSON序列化等配置
class Config:
# 定义datetime和date类型的JSON编码器,将其格式化为指定字符串
json_encoders = {
# 若为datetime类型,格式化为“年-月-日 时:分:秒”,若为None则保持None
datetime: lambda dt: dt.strftime("%Y-%m-%d %H:%M:%S") if dt is not None else None,
# 若为date类型,格式化为“年-月-日”,若为None则保持None
date: lambda dt: dt.strftime("%Y-%m-%d") if dt is not None else None
}
# 定义模型验证器,在数据解析前(mode='before')执行,用于处理字符串格式的日期时间
@model_validator(mode='before')
def parse_string_datetimes(cls, data: dict) -> dict:
# 处理datetime类型字段:将字符串格式的日期时间转换为datetime对象
datetime_fields = {
k: datetime.strptime(v, "%Y-%m-%d %H:%M:%S") # 使用strptime解析字符串为datetime
for k, v in data.items() # 遍历输入数据的键值对
if isinstance(v, str) # 只处理值为字符串的项
and k in cls.model_fields # 键必须是模型中定义的字段
and cls.model_fields[k].annotation is datetime # 字段的注解类型是datetime
}
# 处理date类型字段:将字符串格式的日期转换为date对象(通过datetime解析后取date部分)
date_fields = {
k: datetime.strptime(v, "%Y-%m-%d").date() # 使用strptime解析字符串为datetime后取date
for k, v in data.items() # 遍历输入数据的键值对
if isinstance(v, str) # 只处理值为字符串的项
and k in cls.model_fields # 键必须是模型中定义的字段
and cls.model_fields[k].annotation is date # 字段的注解类型是date
}
# 打印转换后的datetime字段,用于调试
# print("datetime_fields", datetime_fields)
# 合并原始数据、转换后的datetime字段和date字段,转换后的字段会覆盖原始数据中的对应键
result = {**data, **datetime_fields, **date_fields}
# 打印合并后的结果,用于调试
# print("result", result)
# 返回处理后的数据字典
return result
+15
View File
@@ -0,0 +1,15 @@
from sqlmodel import Field
from app.model.BasicModel import BasicModel
from app.utils.create_module_service import create_model_service
class LgApprove(BasicModel, table=True):
__tablename__ = "lg_approve"
status: str = Field(default=None, description="报销单的状态")
remarks: str = Field(default=None, description="报销备注信息")
result_content: str = Field(default=None, description="审批结果信息")
LgApproveService = create_model_service(LgApprove)
+13
View File
@@ -0,0 +1,13 @@
from sqlmodel import Field
from app.model.BasicModel import BasicModel
from app.utils.create_module_service import create_model_service
class LgChat(BasicModel, table=True):
__tablename__ = "lg_chat"
title: str = Field(default=None, description="会话标题")
LgChatService = create_model_service(LgChat)
+16
View File
@@ -0,0 +1,16 @@
from sqlmodel import Field
from app.model.BasicModel import BasicModel
from app.utils.create_module_service import create_model_service
class LgMessage(BasicModel, table=True):
__tablename__ = "lg_message"
title: str = Field(default=None, description="消息标题")
content: str = Field(default=None, description="消息内容")
status: str = Field(default=None, description="消息的状态")
render_configs: str = Field(default=None, description="渲染配置")
LgMessageService = create_model_service(LgMessage)
+16
View File
@@ -0,0 +1,16 @@
from datetime import datetime
from sqlmodel import Field
from app.model.BasicModel import BasicModel
from app.utils.create_module_service import create_model_service
class LlmOrder(BasicModel, table=True):
__tablename__ = "llm_order"
prod_id: str = Field(default=None, description="商品ID")
user_id: str = Field(default=None, description="用户ID")
LlmOrderService = create_model_service(LlmOrder)
+18
View File
@@ -0,0 +1,18 @@
from datetime import datetime
from sqlmodel import Field
from app.model.BasicModel import BasicModel
from app.utils.create_module_service import create_model_service
class LlmProduct(BasicModel, table=True):
__tablename__ = "llm_product"
name: str = Field(default=None, description="商品名称")
price: float = Field(default=None, description="商品价格")
valid_start: datetime = Field(default=None, description="商品有效开始时间")
valid_end: datetime = Field(default=None, description="商品有效结束时间")
LlmProductService = create_model_service(LlmProduct)
+15
View File
@@ -0,0 +1,15 @@
from datetime import datetime
from sqlmodel import Field
from app.model.BasicModel import BasicModel
class LlmUser(BasicModel, table=True):
__tablename__ = "llm_user"
full_name: str = Field(default=None, description="用户名称")
username: str = Field(default=None, description="用户名")
password: str = Field(default=None, description="用户密码")
member_start: datetime = Field(default=None, description="开通会员时间")
member_end: datetime = Field(default=None, description="会员截止到期时间")
+44
View File
@@ -0,0 +1,44 @@
import socket
import psutil
from app.config.env import env
def run_uvicorn():
import uvicorn
# 获取环境变量中的端口号
port = int(env.server_port)
# 打印所有可用的访问地址
print("\n服务已启动,以下是可用的访问地址:")
print(f" - 本地访问: http://127.0.0.1:{port}")
for ip in get_local_ips():
print(f" - 网络访问: http://{ip}:{port}")
print("\n") # 空行美化输出
"""
uvicorn.run() 启动了一个异步事件循环来处理 HTTP 请求,
这个循环会一直运行直到服务器被手动停止(比如按下 Ctrl+C)。
因此,uvicorn.run() 之后的代码不会被执行,直到服务器关闭。
"""
# 启动Uvicorn服务器
uvicorn.run("app.server:app", host="0.0.0.0", port=port)
# uvicorn.run之后的代码永远都不会执行
# 使用FastAPI的 @app.on_event("startup")装饰器可以在服务器成功启动后执行代码
print('App is running...(Never Callable)')
def get_local_ips():
ips = []
try:
for interface, addrs in psutil.net_if_addrs().items():
for addr in addrs:
if addr.family == socket.AF_INET and addr.address != '127.0.0.1':
ips.append(addr.address)
except Exception as e:
print(f"获取 IP 地址时出错: {e}")
return ips
+88
View File
@@ -0,0 +1,88 @@
from fastapi import Query
from langchain_core.messages import HumanMessage
from langchain_core.output_parsers import StrOutputParser
from langchain_core.runnables import RunnableLambda
from langserve import add_routes
from app.config.env import env
from app.controller.add_langgraph_approve_route import add_langgraph_approve_route
from app.controller.add_langgraph_chat_route import add_langgraph_chat_route
from app.controller.add_langgraph_route import add_langgraph_route
from app.controller.add_lg_approve_route import add_lg_approve_route
from app.controller.add_sqlmodel_route import add_sqlmodel_route
from app.controller.add_user_route import add_user_route
from app.controller.custom_chat_playground import add_custom_chat_playground_route
from app.controller.custom_stream_api import add_custom_stream_api_route
from app.controller.test_connection import add_test_connection_route
from app.controller.test_sqlmodel import add_test_sqlmodel_route
from app.controller.test_sync import add_test_sync_route
from app.controller.translate_controller import add_translate_route
from app.create_app import create_app
from app.middlewares.app_middlewares import add_app_middlewares
from app.model.LgApprove import LgApproveService
from app.model.LgChat import LgChatService
from app.model.LgMessage import LgMessageService
from app.model.LlmOrder import LlmOrder, LlmOrderService
from app.model.LlmProduct import LlmProduct, LlmProductService
from app.run_uvicorn import run_uvicorn
from app.utils.ModelInputSchema import ModelInputSchema
from app.utils.add_async_route import add_async_route
from app.utils.create_module_service import create_model_service
from app.utils.llm_utils import create_llm
from app.utils.next_id import add_next_id_route
app = create_app()
add_app_middlewares(app)
add_translate_route(app)
add_custom_chat_playground_route(app)
add_test_sync_route(app)
add_custom_stream_api_route(app)
add_test_connection_route(app)
add_test_sqlmodel_route(app)
add_next_id_route(app)
add_sqlmodel_route(app)
add_user_route(app)
add_langgraph_route(app)
add_lg_approve_route(app)
add_langgraph_approve_route(app)
add_langgraph_chat_route(app)
@app.get("/get_env")
async def test():
return env.model_dump_json()
@app.get("/test")
async def test():
return {"msg": "hello"}
@app.get("/test_llm")
async def test_llm(user_content: str = Query(..., description="用户输入的文本内容,将传递给大语言模型处理")):
return (create_llm() | StrOutputParser()).invoke([HumanMessage(content=user_content)])
add_routes(
app=app,
runnable=RunnableLambda(lambda x: x['messages']) | create_llm() | StrOutputParser(),
input_type=ModelInputSchema,
path="/doubao"
)
add_async_route(
app=app,
runnable=RunnableLambda(lambda x: x['messages']) | create_llm("bailian-qwen-turbo").with_types(input_type=ModelInputSchema),
path="/qwen"
)
LlmOrderService.add_route(app=app, path="/llm_order")
LlmProductService.add_route(app=app, path="/llm_product")
LgApproveService.add_route(app=app, path="/lg_approve")
LgMessageService.add_route(app=app, path="/lg_message")
LgChatService.add_route(app=app, path="/lg_chat")
if __name__ == "__main__":
run_uvicorn()
+34
View File
@@ -0,0 +1,34 @@
from datetime import timedelta, datetime, timezone
import jwt
from passlib.context import CryptContext
from app.config.env import env
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
class CryptUtils:
@staticmethod
def get_password_hash(password: str):
return pwd_context.hash(password)
@staticmethod
def verify_password(plain_password: str, hashed_password: str):
return pwd_context.verify(plain_password, hashed_password)
@staticmethod
def create_access_token(username: str, expires_delta: timedelta | None = None):
data: dict = {"sub": username}
if expires_delta:
expire = datetime.now(timezone.utc) + expires_delta
else:
expire = datetime.now(timezone.utc) + timedelta(minutes=env.jwt_access_token_expire_minutes)
data.update({'exp': expire})
return jwt.encode(data, env.jwt_secret_key, env.jwt_algorithm)
@staticmethod
def get_username_from_token(token: str):
data = jwt.decode(token, env.jwt_secret_key, algorithms=[env.jwt_algorithm])
username = data.get("sub")
return username
+72
View File
@@ -0,0 +1,72 @@
import json
import requests
from langchain_core.embeddings import Embeddings
class CustomEmbeddings(Embeddings):
"""
自定义文本嵌入类,用于将文本转换为向量表示
继承自LangChain的Embeddings基类
"""
def __init__(self, base_url, api_key, model):
"""
初始化自定义嵌入类
参数:
base_url: API的基础URL
api_key: 访问API所需的密钥
model: 要使用的嵌入模型名称
"""
self.base_url = base_url # API基础URL
self.api_key = api_key # API访问密钥
self.model = model # 嵌入模型名称
def embed_documents(self, texts):
"""
将多个文档转换为嵌入向量
参数:
texts: 包含多个文本的列表
返回:
包含每个文本对应嵌入向量的列表
"""
# 设置请求头,包括内容类型和认证信息
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}"
}
# 构建请求负载
payload = {"input": texts, "model": self.model, "encoding_format": "float"}
# 发送POST请求到嵌入API
response = requests.post(
f"{self.base_url}/embeddings",
headers=headers,
data=json.dumps(payload)
)
# 检查请求是否成功,如果失败则抛出异常
response.raise_for_status()
# 解析响应JSON数据
json_data = response.json()
# 从响应数据中提取嵌入向量并返回
return [item["embedding"] for item in json_data["data"]]
def embed_query(self, text):
"""
将单个查询文本转换为嵌入向量
参数:
text: 查询文本
返回:
对应的嵌入向量
"""
# 调用embed_documents处理单个文本,并返回第一个结果
return self.embed_documents([text])[0]
+13
View File
@@ -0,0 +1,13 @@
from typing import List, Union
from langchain_core.messages import HumanMessage, AIMessage, SystemMessage
from pydantic import BaseModel, Field
class ModelInputSchema(BaseModel):
"""Input for the chat endpoint."""
messages: List[Union[HumanMessage, AIMessage, SystemMessage]] = Field(
...,
description="当前对话中的消息历史",
)
+14
View File
@@ -0,0 +1,14 @@
from pydantic import BaseModel, Field
class PageQueryParams(BaseModel):
page: int = Field(default=0, description="分页查询的页数")
page_size: int = Field(default=5, description="分页查询每页条数")
all: bool = Field(default=False, description="是否查询所有数据,也就是不分页")
count: bool = Field(default=True, description="是否查询总数")
sort_field: str = Field(default="created_at", description="排序字段")
sort_desc: str = Field(default="desc", description="排序方式")
filters: dict = Field(default=None, description="筛选参数")
+70
View File
@@ -0,0 +1,70 @@
import json
import time
from typing import Any, List
from fastapi import APIRouter, FastAPI
from langchain_core.messages import AIMessage
from langchain_core.runnables import Runnable
from starlette.responses import StreamingResponse
def format_ai_message(ai_message: AIMessage):
return {
"choices": [{
"finish_reason": "stop",
"index": 0,
"message": {
"content": ai_message.content,
"role": "assistant"
}
}],
"created": int(time.time()),
"id": ai_message.id,
"usage": ai_message.response_metadata.get('token_usage')
}
def add_async_route(
app: FastAPI,
runnable: Runnable,
path: str,
input_type: Any = None, # 设置接口的传入参数类型,input_type实际上是一个pydantic类
):
router = APIRouter(prefix=path, tags=[path])
_input_type = input_type or runnable.input_schema
@router.post('/ainvoke')
async def ainvoke(input: _input_type):
ai_message = await runnable.ainvoke(input.model_dump())
return format_ai_message(ai_message)
@router.post('/abatch')
async def abatch(inputs: List[_input_type]):
ai_message_list = await runnable.abatch([input.model_dump() for input in inputs])
return [format_ai_message(ai_message) for ai_message in ai_message_list]
@router.post('/astream')
async def astream(input: _input_type):
async def generator_function():
result_template = {
"choices": [{"delta": {}, "index": 0}],
"created": time.time(),
"id": "",
"usage": None
}
async for chunk in runnable.astream(input.model_dump()):
result_template["choices"][0]["delta"]['content'] = chunk.content
result_template["choices"][0]["delta"]['role'] = 'assistant'
if chunk.response_metadata.get('finish_reason') is not None:
result_template["choices"][0]["delta"]['finish_reason'] = chunk.response_metadata.get('finish_reason')
result_template['id'] = chunk.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.include_router(router)
+458
View File
@@ -0,0 +1,458 @@
import asyncio
import json
from typing import Type, List, Any, Union
from fastapi import FastAPI, APIRouter, HTTPException, Body
from pydantic import create_model
from sqlalchemy import func
from sqlmodel import select
from app.model.BasicModel import BasicModel
from app.utils.PageQueryParams import PageQueryParams
from app.utils.db_utils import AsyncSessionDep
from app.utils.next_id import next_id
def create_model_service(
#/*@formatter:off*/
Cls: Type[BasicModel], # model实体类
before_query_list=None, # 分页查询前异步处理函数,参数:(query_param, session)
after_query_list=None, # 分页查询后异步处理函数,参数:(query_cls_list, has_next, query_param, session)
before_query_item=None, # 单条查询前异步处理函数,参数:(row_dict, session)
after_query_item=None, # 单条查询后异步处理函数,参数:(item_cls, row_dict, session)
before_insert=None, # 单条新建前异步处理函数,参数:(row_dict, session)
after_insert=None, # 单条新建后异步处理函数,参数:(insert_cls, row_dict, session)
before_update=None, # 单条更新前异步处理函数,参数:(row_dict, session)
after_update=None, # 单条更新后异步处理函数,参数:(update_cls, row_dict, session)
before_delete=None, # 单条删除前异步处理函数,参数:(row_dict, session)
after_delete=None, # 单条删除后异步处理函数,参数:(delete_cls, row_dict, session)
before_batch_insert=None, # 批量新建前异步处理函数,参数:(row_dict_list, session)
after_batch_insert=None, # 批量新建后异步处理函数,参数:(refresh_cls_list, row_dict_list, session)
before_batch_update=None, # 批量更新前异步处理函数,参数:(row_dict_list, session)
after_batch_update=None, # 批量更新后异步处理函数,参数:(refresh_cls_list, row_dict_list, session)
before_batch_delete=None, # 批量删除前异步处理函数,参数:(row_dict_list, session)
after_batch_delete=None, # 批量删除后异步处理函数,参数:(delete_cls_list, row_dict_list, session)
# /*@formatter:on*/
):
# 定义模型服务类,封装模型相关的CRUD接口及业务逻辑
class ModelService:
# 支持的所有端点列表,包含常用的CRUD及批量操作
END_POINTS = ['list', 'item', 'insert', 'batch_insert', 'update', 'batch_update', 'delete', 'batch_delete']
def __init__(self):
# 验证传入的模型类是否继承自BasicModel,确保基础字段存在
if not issubclass(Cls, BasicModel):
raise TypeError(f"{Cls.__name__} 必须继承自 BasicModel")
# 保存当前操作的模型类
self.Cls = Cls
# 检查字典中的键是否为模型类的有效属性
# 参数:
# row_dict: 待检查的字典(通常为请求参数)
def check_invalid_keys(self, row_dict: dict):
# 筛选出所有不在模型类属性中的键(无效键)
invalid_keys = [key for key in row_dict.keys() if not hasattr(Cls, key)]
if invalid_keys:
# 若存在无效键,抛出HTTP 500异常,提示无效键和有效键列表
raise HTTPException(
status_code=500,
detail=f"Invalid filter keys: {invalid_keys}. Valid keys are: {Cls.__annotations__.keys()}"
)
def add_route(
self,
app: FastAPI, # FastAPI实例,用来注册路由
path: str, # 路由前缀地址
end_points: List[str] = None, # 生成的端点入口接口清单
):
# 确定启用的端点,默认为全部支持的端点
if not end_points:
end_points = self.END_POINTS
# 动态创建分页查询的响应模型:包含数据列表和是否有下一页的标识
ListResponse = create_model(f"{Cls.__name__}ListResponse", list=(List[Cls], ...), has_next=(bool, ...), total=(Union[int, None], None))
# 动态创建单条查询的响应模型:包含单个模型实例
ItemResponse = create_model(f"{Cls.__name__}ItemResponse", result=(Cls, ...))
# 动态创建批量操作的响应模型:包含操作后的模型实例列表
BatchResponse = create_model(f"{Cls.__name__}BatchResponse", result=(List[Cls], ...))
# 动态创建批量删除的响应模型:包含删除操作是否成功的标识
DeleteResponse = create_model(f"{Cls.__name__}BatchResponse", result=(bool, ...))
# 创建APIRouter实例,设置路由前缀和标签(标签用于API文档分组)
router = APIRouter(prefix=path, tags=[path])
# 若启用"list"端点,注册列表查询接口
if 'list' in end_points:
# 列表查询接口:支持过滤和分页,响应模型为ListResponse
@router.post("/list", response_model=ListResponse)
async def _list(query_param: PageQueryParams, session: AsyncSessionDep):
# 调用query_list方法执行查询,获取数据列表和是否有下一页
query_cls_list, has_next, total = await self.query_list(query_param, session)
# 返回符合响应模型的结果
return {
"list": query_cls_list,
"has_next": has_next,
"total": total,
}
# 若启用"item"端点,注册单条查询接口
if 'item' in end_points:
# 单条查询接口:根据条件查询单条记录,响应模型为ItemResponse
@router.post("/item", response_model=ItemResponse)
async def _item(
session: AsyncSessionDep,
row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}")
):
# 调用query_item方法查询单条记录并返回
return {"result": await self.query_item(session, row_dict)}
# 若启用"insert"端点,注册单条插入接口
if 'insert' in end_points:
# 单条插入接口:新增一条记录,响应模型为ItemResponse
@router.post("/insert", response_model=ItemResponse)
async def _insert(
session: AsyncSessionDep,
row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}")
):
# 调用item_insert方法执行插入并返回结果
return {"result": await self.item_insert(session, row_dict)}
# 若启用"batch_insert"端点,注册批量插入接口
if 'batch_insert' in end_points:
# 批量插入接口:批量新增记录,响应模型为BatchResponse
@router.post("/batch_insert", response_model=BatchResponse)
async def _batch_insert(
session: AsyncSessionDep,
row_dict_list: List[dict] = Body(..., description=f"批量插入的数据数组,字段参考{Cls.__name__}")
):
# 调用batch_insert方法执行批量插入并返回结果
return {"result": await self.batch_insert(session, row_dict_list)}
# 若启用"update"端点,注册单条更新接口
if 'update' in end_points:
# 单条更新接口:更新一条记录,响应模型为ItemResponse
@router.post("/update", response_model=ItemResponse)
async def _update(
session: AsyncSessionDep,
row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}")
):
# 调用item_update方法执行更新并返回结果
return {"result": await self.item_update(session, row_dict)}
# 若启用"batch_update"端点,注册批量更新接口
if 'batch_update' in end_points:
# 批量更新接口:批量更新记录,响应模型为BatchResponse
@router.post("/batch_update", response_model=BatchResponse)
async def _batch_update(
session: AsyncSessionDep,
row_dict_list: List[dict] = Body(..., description=f"批量更新的数据数组,字段参考{Cls.__name__}")
):
# 调用batch_update方法执行批量更新并返回结果
return {"result": await self.batch_update(session, row_dict_list)}
# 若启用"delete"端点,注册单条删除接口
if 'delete' in end_points:
# 单条删除接口:删除一条记录,响应模型为DeleteResponse
@router.post("/delete", response_model=DeleteResponse)
async def _delete(
session: AsyncSessionDep,
row_dict: dict = Body(..., description=f"删除的数据,字段参考{Cls.__name__}")
):
# 调用item_delete方法执行删除并返回结果
return {"result": await self.item_delete(session, row_dict)}
# 若启用"batch_delete"端点,注册批量删除接口
if 'batch_delete' in end_points:
# 批量删除接口:批量删除记录,响应模型为DeleteResponse
@router.post("/batch_delete", response_model=DeleteResponse)
async def _delete(
session: AsyncSessionDep,
row_dict_list: List[dict] = Body(..., description=f"批量删除的数据数组,字段参考{Cls.__name__}")
):
# 调用batch_delete方法执行批量删除并返回结果
return {"result": await self.batch_delete(session, row_dict_list)}
# 将路由添加到FastAPI应用
app.include_router(router)
# 分页查询工具方法:执行带过滤和分页的查询
async def query_list(self, query_param: PageQueryParams, session: AsyncSessionDep):
if before_query_list is not None:
await before_query_list(query_param, session)
# 创建基础查询:查询当前模型类的所有记录
query = select(Cls)
count_query = select(func.count()).select_from(Cls)
# 若有过滤条件,验证并应用过滤
if query_param.filters:
self.check_invalid_keys(query_param.filters)
# 为每个过滤条件添加WHERE子句(字段=值)
for key, value in query_param.filters.items():
query = query.where(getattr(Cls, key) == value)
count_query = count_query.where(getattr(Cls, key) == value)
if query_param.sort_field:
# 为排序字段添加ORDER BY子句
cls_attr = getattr(Cls, query_param.sort_field)
order_value = cls_attr.desc() if query_param.sort_desc == 'desc' else cls_attr.asc();
query = query.order_by(order_value)
# 若不查询全部数据(即启用分页)
if query_param.all is False:
# 计算偏移量(跳过前N条),并查询比一页多1条的记录(用于判断是否有下一页)
query = query.offset(query_param.page * query_param.page_size).limit(query_param.page_size + 1)
# 执行查询并获取结果
result = await session.execute(query)
if query_param.count:
total, = (await session.execute(count_query)).one()
else:
total = None
# 将查询结果转换为标量列表(模型实例列表)
query_cls_list: List[Any] = result.scalars().all()
# 打印查询结果类型和内容(调试用)
print("query_cls_list", type(query_cls_list), query_cls_list)
# 判断是否有下一页:若查询结果数量等于一页大小+1,则说明有下一页
has_next = len(query_cls_list) == query_param.page_size + 1
# 若有下一页,移除多查询的那一条记录
if has_next:
query_cls_list.pop()
if after_query_list is not None:
await after_query_list(query_cls_list, has_next, query_param, session)
# 返回处理后的结果列表和是否有下一页的标识
return (query_cls_list, has_next, total)
# 单条查询工具方法:根据条件查询单条记录
async def query_item(self, session: AsyncSessionDep, row_dict: dict = Body(..., description=f"查询数据的字段筛选值,字段参考{Cls.__name__}")):
if before_query_item is not None:
await before_query_item(row_dict, session)
# 创建基础查询:查询当前模型类的所有记录
query = select(Cls)
# 验证查询条件中的键是否有效
self.check_invalid_keys(row_dict)
# 为每个条件添加WHERE子句(字段=值)
for key, value in row_dict.items():
query = query.where(getattr(Cls, key) == value)
# 执行查询
result = await session.execute(query)
# 返回第一条匹配的记录(若存在)
item_cls = result.scalars().first()
if after_query_item is not None:
await after_query_item(item_cls, row_dict, session)
return item_cls
# 单条插入工具方法:新增一条记录
async def item_insert(self, session: AsyncSessionDep, row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}")):
if before_insert is not None:
await before_insert(row_dict, session)
# 若未提供id,自动生成唯一id
if row_dict.get("id") is None:
row_dict["id"] = await next_id()
try:
# 使用模型类验证数据并创建实例(校验字段类型和约束)
insert_cls = Cls.model_validate(row_dict)
except ValueError as e:
# 数据验证失败时,抛出HTTP 500异常并返回错误详情
raise HTTPException(status_code=500, detail=str(e))
# 将实例添加到数据库会话
session.add(insert_cls)
# 提交事务(保存到数据库)
await session.commit()
# 刷新实例,获取数据库生成的最新数据(如自动更新的时间字段)
await session.refresh(insert_cls)
if after_insert is not None:
await after_insert(insert_cls, row_dict, session)
# 返回插入的实例
return insert_cls
# 批量插入工具方法:批量新增记录
async def batch_insert(self, session: AsyncSessionDep, row_dict_list: List[dict] = Body(..., description=f"批量插入的数据数组,字段参考{Cls.__name__}")):
if before_batch_insert is not None:
await before_batch_insert(row_dict_list, session)
# 筛选出没有id的记录(需要自动生成id)
row_dict_list_without_id = []
for row_dict in row_dict_list:
if row_dict.get("id") is None:
row_dict_list_without_id.append(row_dict)
# 若存在需要自动生成id的记录
if len(row_dict_list_without_id):
# 批量生成唯一id(数量等于需要生成id的记录数)
new_id_list = await next_id(len(row_dict_list_without_id))
# 为每条记录分配生成的id
for index, id in enumerate(new_id_list):
row_dict_list_without_id[index]["id"] = id
try:
# 验证所有记录并转换为模型实例列表
insert_cls_list = [Cls.model_validate(row_dict) for row_dict in row_dict_list]
except ValueError as e:
# 验证失败时抛出异常
raise HTTPException(status_code=500, detail=str(e))
# 将所有实例添加到会话
session.add_all(insert_cls_list)
# 提交事务,保存数据到数据库
await session.commit()
# 查询并返回所有插入的实例(刷新数据,确保获取最新状态)
refresh_cls_list = (await session.execute(select(Cls).where(Cls.id.in_([obj.id for obj in insert_cls_list])))).scalars().all()
if after_batch_insert is not None:
await after_batch_insert(refresh_cls_list, row_dict_list, session)
# 返回刷新后的实例列表
return refresh_cls_list
# 单条更新工具方法:更新一条记录
async def item_update(self, session: AsyncSessionDep, row_dict: dict = Body(..., description=f"更新的数据,字段参考{Cls.__name__}")):
if before_update is not None:
await before_update(row_dict, session)
# 检查id是否存在(更新必须指定id)
if not row_dict.get('id'):
raise HTTPException(status_code=400, detail="ID不能为空")
# 根据id查询要更新的记录
update_cls = (await session.exec(select(Cls).where(Cls.id == row_dict.get('id')))).first()
if not update_cls:
# 若记录不存在,抛出异常
raise HTTPException(status_code=500, detail="Update row not found")
# 遍历更新字段:为记录的每个键设置新值
for key, value in row_dict.items():
setattr(update_cls, key, value)
# 将更新后的实例添加到会话
session.add(update_cls)
# 提交事务
await session.commit()
# 刷新实例,获取最新数据
await session.refresh(update_cls)
if after_update is not None:
await after_update(update_cls, row_dict, session)
# 返回更新后的实例
return update_cls
# 批量更新工具方法:批量更新记录
async def batch_update(self, session: AsyncSessionDep, row_dict_list: List[dict] = Body(..., description=f"批量更新的数据数组,字段参考{Cls.__name__}")):
if before_batch_update is not None:
await before_batch_update(row_dict_list, session)
# 提取所有待更新记录的id
update_id_list = [row_dict['id'] for row_dict in row_dict_list]
# 根据id查询所有待更新的记录
update_cls_list = (await session.exec(select(Cls).where(Cls.id.in_(update_id_list)))).all()
# 若查询到的记录数量与待更新数量不一致,说明部分id不存在
if len(update_cls_list) != len(row_dict_list):
# 抛出异常并提示不存在的id
raise HTTPException(status_code=500, detail="Update row not found:" + json.dumps(row_dict_list, ensure_ascii=False))
# 创建id到更新数据的映射(便于快速查找)
id_2_row_dict = {row_dict["id"]: row_dict for row_dict in row_dict_list}
# 遍历每条查询到的记录,更新其字段
for update_cls in update_cls_list:
# 获取当前记录对应的更新数据(根据id)
row_dict = id_2_row_dict.get(update_cls.id, None)
# 遍历更新字段
for key, value in row_dict.items():
setattr(update_cls, key, value)
# 将所有更新后的实例添加到会话
session.add_all(update_cls_list)
# 提交事务
await session.commit()
# 查询并返回所有更新后的实例(刷新数据)
refresh_cls_list = (await session.execute(select(Cls).where(Cls.id.in_([obj.id for obj in update_cls_list])))).scalars().all()
if after_batch_update is not None:
await after_batch_update(refresh_cls_list, row_dict_list, session)
# 返回刷新后的实例列表
return refresh_cls_list
# 单条删除工具方法:删除一条记录
async def item_delete(self, session: AsyncSessionDep, row_dict: dict = Body(..., description=f"删除的数据,字段参考{Cls.__name__}")):
if before_delete is not None:
await before_delete(row_dict, session)
# 根据id查询要删除的记录
delete_cls = (await session.exec(select(Cls).where(Cls.id == row_dict.get('id')))).first()
if not delete_cls:
# 若记录不存在,返回删除失败
return False
# 从会话中删除记录
await session.delete(delete_cls)
# 提交事务,执行删除
await session.commit()
if after_delete is not None:
await after_delete(delete_cls, row_dict, session)
# 返回删除成功
return True
# 批量删除工具方法:批量删除记录
async def batch_delete(self, session: AsyncSessionDep, row_dict_list: List[dict] = Body(..., description=f"批量删除的数据数组,字段参考{Cls.__name__}")):
if before_batch_delete is not None:
await before_batch_delete(row_dict_list, session)
# 若待删除列表为空,返回失败
if not row_dict_list:
return False
# 提取所有待删除记录的id
row_id_list = [row_dict.get("id") for row_dict in row_dict_list]
# 根据id查询所有待删除的记录
delete_cls_list = (await session.exec(select(Cls).where(Cls.id.in_(row_id_list)))).all()
# 若查询到的记录数量与待删除数量不一致,说明部分id不存在
if len(delete_cls_list) != len(row_id_list):
# 抛出异常并提示不存在的id
raise HTTPException(status_code=500, detail="Delete row not found:" + json.dumps(row_id_list, ensure_ascii=False))
# 异步批量删除所有记录(使用gather并发执行删除操作)
await asyncio.gather(*[asyncio.create_task(session.delete(delete_cls)) for delete_cls in delete_cls_list])
# 提交事务,执行删除
await session.commit()
if after_batch_delete is not None:
await after_batch_delete(delete_cls_list, row_dict_list, session)
# 返回删除成功
return True
# 创建并返回ModelService实例
return ModelService()
+62
View File
@@ -0,0 +1,62 @@
import asyncio
import sys
from contextlib import asynccontextmanager
from typing import Annotated, AsyncContextManager
from fastapi.params import Depends
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
from sqlalchemy import AsyncAdaptedQueuePool, text
from sqlalchemy.ext.asyncio import AsyncEngine
from sqlalchemy.orm import sessionmaker
from sqlmodel.ext.asyncio.session import AsyncSession
from app.config.env import env
from sqlmodel import create_engine
DATABASE_URL = f"mysql+asyncmy://{env.db_username}:{env.db_password}@{env.db_host}:{env.db_port}/{env.db_database}?charset=utf8mb4"
# 创建异步引擎实例,用于异步操作数据库
async_engine = AsyncEngine(create_engine(
DATABASE_URL,
poolclass=AsyncAdaptedQueuePool, # 使用异步适配的队列池
pool_size=5, # 连接池保持的连接数
max_overflow=10, # 允许超过pool_size的最大连接数
pool_timeout=30, # 获取连接的超时时间(秒)
pool_recycle=3600, # 连接回收时间(秒)
echo=True, # 启用SQL语句日志输出,便于开发调试
future=True, # 启用SQLAlchemy 2.0风格的未来模式API
))
# 创建一个会话工厂函数
async_session = sessionmaker(
bind=async_engine,
class_=AsyncSession,
expire_on_commit=False
)
# 作用:用于在接口中注入得到会话实例对象session,在接口执行完毕之后,自动执行close动作关闭会话
async def get_async_session() -> AsyncSession:
async with async_session() as session:
yield session
AsyncSessionDep = Annotated[AsyncSession, Depends(get_async_session)]
# 用于启动服务的时候检查数据库连接是否正常
async def check_database_connection():
"""检查数据库连接是否正常"""
try:
async with async_engine.begin() as conn:
print("Connecting Mysql...")
await conn.execute(text("select 1"))
# 打印连接成功信息及连接URL
print("✅ Database connection successful:", DATABASE_URL)
except Exception as e:
# 打印连接失败信息及错误详情
print(f"❌ Database connection failed: {e}")
# 重新抛出异常,让上层处理
raise e
return async_engine
+62
View File
@@ -0,0 +1,62 @@
from langchain_core.runnables import RunnableLambda
from langchain_openai import ChatOpenAI
from app.config.ai_configs import ai_configs
from app.utils.CustomEmbeddings import CustomEmbeddings
def create_llm(platform_code='huoshan-doubao', temperature=0.5):
_ai_config = ai_configs.get(platform_code)
if _ai_config is None:
raise Exception('Unknown platform code', platform_code)
return ChatOpenAI(
base_url=_ai_config.get('url').replace("chat/completions", ""),
api_key=_ai_config.get('key'),
model=_ai_config.get('model'),
temperature=temperature,
)
def create_embeddings(platform_code="huoshan-embedding-240715"):
"""
创建自定义嵌入模型实例
参数:
platform_code: 平台代码,用于从默认配置中查找对应平台的API信息
返回:
CustomEmbeddings类的实例,用于生成文本嵌入向量
异常:
当找不到对应平台代码的配置时抛出异常
"""
# 从默认配置中获取指定平台的AI配置信息
_ai_config = ai_configs.get(platform_code)
# 检查配置是否存在
if _ai_config is None:
raise Exception('Unknown platform code', platform_code)
# 创建并返回自定义嵌入模型实例
return CustomEmbeddings(
base_url=_ai_config.get('url').replace("/embeddings", ""), # API基础URL
api_key=_ai_config.get('key'), # API密钥
model=_ai_config.get('model') # 嵌入模型名称
)
def chain_log(format_func=None):
"""创建一个函数,用于在链中打印上一个管道的结果"""
def func(val):
print("\033[34m chain log==>>", format_func(val) if format_func is not None else val, '\033[0m')
return val
return func
def runnable_chain_log(format_func=None):
"""创建一个Runnable对象,用于在链中打印上一个管道的结果,如果上一个管道是字典对象,那么打印这个字典对象需要使用runnable_chain_log"""
return RunnableLambda(chain_log(format_func))
+23
View File
@@ -0,0 +1,23 @@
from fastapi import FastAPI
from sqlalchemy import text
from app.utils.db_utils import async_session
async def next_id(num: int = 1):
async with async_session() as session:
sql_string = "select " + ",".join([f"uuid() as _{index}" for index in range(num)])
print(sql_string)
result = await session.execute(text(sql_string))
val = result.first()
print(val)
arr = list(val or [])
return arr[0] if num == 1 else arr
def add_next_id_route(app: FastAPI):
@app.get("/next_id")
async def _next_id(num: int = 1):
return {
"data": await next_id(num),
}
+136
View File
@@ -0,0 +1,136 @@
import asyncio
import sys
import time
from typing import Optional, AsyncContextManager, Annotated, TypedDict
from fastapi import Depends
from langgraph.checkpoint.postgres.aio import AsyncPostgresSaver
from langgraph.graph import StateGraph
from langgraph.graph.state import CompiledStateGraph
from app.config.env import env
# 构造PostgreSQL数据库连接字符串
POSTGRES_DATABASE_URL = (f"postgresql://{env.pg_db_username}:{env.pg_db_password}@"
f"{env.pg_db_host}:{env.pg_db_port}/{env.pg_db_database}"
f"?connect_timeout=10&keepalives=1&keepalives_idle=30"
f"&keepalives_interval=10&keepalives_count=3")
class PostgresCheckpointerManager:
# 单例实例,存储AsyncPostgresSaver对象
_instance: Optional[AsyncPostgresSaver] = None
# 存储异步上下文管理器,用于正确管理数据库连接的生命周期
_context_manger: Optional[AsyncContextManager] = None
# 异步锁,确保在并发环境下单例实例的创建是线程安全的
_lock = asyncio.Lock()
# 最后一次检测连接是否有效的时间
_last_check_time = time.time()
# 一个图用来测试连接是否仍然有效
_graph: Optional[CompiledStateGraph] = None
@staticmethod
async def get_instance() -> AsyncPostgresSaver:
is_connection_alive = await PostgresCheckpointerManager.is_connection_alive()
if not is_connection_alive:
# 使用异步锁确保并发安全
async with PostgresCheckpointerManager._lock:
if PostgresCheckpointerManager._instance:
await PostgresCheckpointerManager.close_instance()
# 从连接字符串创建AsyncPostgresSaver上下文管理器
PostgresCheckpointerManager._context_manger = AsyncPostgresSaver.from_conn_string(POSTGRES_DATABASE_URL)
# 进入异步上下文,初始化数据库连接
PostgresCheckpointerManager._instance = await PostgresCheckpointerManager._context_manger.__aenter__()
# 创建用来测试连接是否有效的图
PostgresCheckpointerManager._graph = create_test_graph(PostgresCheckpointerManager._instance)
# 打印调试信息
print("Create AsyncPostgresSaver:", PostgresCheckpointerManager._instance)
# 返回单例实例
return PostgresCheckpointerManager._instance
@staticmethod
async def close_instance():
"""
关闭并清理单例实例和相关资源
"""
# 清理实例引用
if PostgresCheckpointerManager._instance is not None:
PostgresCheckpointerManager._instance = None
# 退出上下文管理器,正确关闭数据库连接
if PostgresCheckpointerManager._context_manger is not None:
await PostgresCheckpointerManager._context_manger.__aexit__(None, None, None)
PostgresCheckpointerManager._context_manger = None
# 清理掉测试连接的图
PostgresCheckpointerManager._graph = None
return
@staticmethod
async def is_connection_alive() -> bool:
"""检查数据库连接是否仍然存活"""
if PostgresCheckpointerManager._instance is None:
return False
# 检测间隔小于60秒,直接返回True
if time.time() - PostgresCheckpointerManager._last_check_time < 60:
return True
try:
print("\n\nCheck Postgres connection...\n\n")
await PostgresCheckpointerManager._graph.aget_state(config={"configurable": {"thread_id": "@@TestAsyncPostgresSaverConnectionIsKeepAlive"}})
PostgresCheckpointerManager._last_check_time = time.time()
return True
except Exception as e:
print(f"Postgres connection check failed: {e}")
return False
# 定义依赖注入类型,用于FastAPI自动注入AsyncPostgresSaver实例
AsyncPostgresSaverDep = Annotated[AsyncPostgresSaver, Depends(PostgresCheckpointerManager.get_instance)]
async def check_postgres_connection():
"""
用于启动服务的时候检查Postgres数据库连接是否正常
"""
try:
print("Connecting Postgres...")
# 尝试获取数据库连接实例
await PostgresCheckpointerManager.get_instance()
# 连接成功,打印成功信息
print("✅ Postgres connection successful:", POSTGRES_DATABASE_URL)
except Exception as e:
# 打印连接失败信息及错误详情
print(f"❌ Postgres connection failed: {e}")
# 重新抛出异常,让上层处理
raise e
async def close_postgres_connection():
await PostgresCheckpointerManager.close_instance()
# 在Windows平台上设置事件循环策略为WindowsSelectorEventLoopPolicy
if sys.platform == "win32":
from asyncio import WindowsSelectorEventLoopPolicy
asyncio.set_event_loop_policy(WindowsSelectorEventLoopPolicy())
def create_test_graph(checkpointer: AsyncPostgresSaver):
class StateSchema(TypedDict):
input: str
builder = StateGraph(StateSchema)
def node(state: StateSchema):
return {}
builder.add_node(node)
builder.set_entry_point("node")
builder.set_finish_point("node")
return builder.compile(checkpointer=checkpointer)
View File
Generated
+2688
View File
File diff suppressed because it is too large Load Diff
+46
View File
@@ -0,0 +1,46 @@
[tool.poetry]
name = "ai-langserve"
version = "0.1.0"
description = ""
authors = ["Your Name <you@example.com>"]
readme = "README.md"
packages = [
{ include = "app" },
]
[tool.poetry.dependencies]
python = "^3.11"
uvicorn = "^0.23.2"
langserve = {extras = ["server"], version = ">=0.0.30"}
pydantic = "2.10.6"
langchain-openai = "^0.3.28"
dotenv = "^0.9.9"
pydantic-settings = "^2.10.1"
psutil = "^7.0.0"
asyncmy = "^0.2.10"
sqlmodel = "^0.0.24"
greenlet = "^3.2.3"
passlib = {extras = ["bcrypt"], version = "^1.7.4"}
pyjwt = "^2.10.1"
python-multipart = "^0.0.20"
langgraph = "^0.6.3"
langgraph-checkpoint-postgres = "^2.0.23"
[tool.poetry.group.dev.dependencies]
langchain-cli = ">=0.0.15"
# 配置国内镜像源
[[tool.poetry.source]]
name = "tsinghua"
url = "https://pypi.tuna.tsinghua.edu.cn/simple"
priority = "primary" # 最高优先级
[[tool.poetry.source]]
name = "aliyun"
url = "https://mirrors.aliyun.com/pypi/simple/"
priority = "supplemental" # 次级优先级
[build-system]
requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
File diff suppressed because one or more lines are too long
+20
View File
File diff suppressed because one or more lines are too long
+1
View File
File diff suppressed because one or more lines are too long