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