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
+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)