feat: general接口拦截器

This commit is contained in:
martsforever
2025-10-15 19:39:13 +08:00
parent b530cb95b7
commit 560c9743dc
7 changed files with 285 additions and 144 deletions
+52
View File
@@ -0,0 +1,52 @@
from typing import List
class GeneralInterceptor():
def __init__(
self,
module: str,
before_list=None,
after_list=None,
before_insert=None,
after_insert=None,
before_update=None,
after_update=None,
before_batch_insert=None,
after_batch_insert=None,
before_batch_update=None,
after_batch_update=None,
before_delete=None,
after_delete=None,
):
self.module = module
self.before_list = before_list
self.after_list = after_list
self.before_insert = before_insert
self.after_insert = after_insert
self.before_update = before_update
self.after_update = after_update
self.before_batch_insert = before_batch_insert
self.after_batch_insert = after_batch_insert
self.before_batch_update = before_batch_update
self.after_batch_update = after_batch_update
self.before_delete = before_delete
self.after_delete = after_delete
_general_interceptors: List[GeneralInterceptor] = []
def add_general_interceptor(interceptor: GeneralInterceptor):
_general_interceptors.append(interceptor)
async def invoke_interceptor(method: str, module: str, *args):
if module.startswith('/'):
module = module[1:]
if module.endswith('/'):
module = module[:-1]
match_interceptors = [item for item in _general_interceptors if item.module == module]
for item in match_interceptors:
interceptor_method = getattr(item, method, None) # 使用 getattr 获取方法
if interceptor_method and callable(interceptor_method):
await interceptor_method(*args) # 调用方法
+28
View File
@@ -1,6 +1,7 @@
import traceback import traceback
from app.config.env import env from app.config.env import env
from app.general.general_interceptors import invoke_interceptor
from app.general.general_utils.build_delete_sql import build_delete_sql from app.general.general_utils.build_delete_sql import build_delete_sql
from app.general.general_utils.build_insert_sql import build_insert_sql from app.general.general_utils.build_insert_sql import build_insert_sql
from app.general.general_utils.build_query_sql import build_query_sql from app.general.general_utils.build_query_sql import build_query_sql
@@ -52,6 +53,8 @@ async def list(session: AsyncSessionDep, query_config, module_config, debug_data
"orders": get_default_orders(query_config, module_config) "orders": get_default_orders(query_config, module_config)
} }
await invoke_interceptor('before_list', module_config['base'], target_query_config, session)
sql, values = build_query_sql(target_query_config, module_config) sql, values = build_query_sql(target_query_config, module_config)
try: try:
@@ -62,6 +65,8 @@ async def list(session: AsyncSessionDep, query_config, module_config, debug_data
create_convertor(module_config)['decode_list'](result) create_convertor(module_config)['decode_list'](result)
await invoke_interceptor('after_list', module_config['base'], result, session)
if n_only_count: if n_only_count:
return { return {
"total": result[0]['total'] "total": result[0]['total']
@@ -120,6 +125,8 @@ async def insert(session: AsyncSessionDep, query_config, module_config, debug_da
row_id = (await get_id(session, 1))[0] row_id = (await get_id(session, 1))[0]
row['id'] = row_id row['id'] = row_id
await invoke_interceptor('before_insert', module_config['base'], row, session)
try: try:
sql, values = build_insert_sql(module_config, row) sql, values = build_insert_sql(module_config, row)
@@ -134,6 +141,8 @@ async def insert(session: AsyncSessionDep, query_config, module_config, debug_da
item_dict = get_value(result, 'result', None) item_dict = get_value(result, 'result', None)
await invoke_interceptor('after_insert', module_config['base'], item_dict, session)
if item_dict is not None: if item_dict is not None:
return { return {
"result": item_dict "result": item_dict
@@ -167,6 +176,8 @@ async def batch_insert(session: AsyncSessionDep, query_config, module_config, de
for index, row in enumerate(rows_without_id): for index, row in enumerate(rows_without_id):
row['id'] = new_id_list[index] row['id'] = new_id_list[index]
await invoke_interceptor('before_batch_insert', module_config['base'], rows, session)
try: try:
for row in rows: for row in rows:
sql, values = build_insert_sql(module_config, row) sql, values = build_insert_sql(module_config, row)
@@ -182,6 +193,9 @@ async def batch_insert(session: AsyncSessionDep, query_config, module_config, de
result = get_value(result, 'list', []) result = get_value(result, 'list', [])
if len(result) > 0: if len(result) > 0:
await invoke_interceptor('after_batch_insert', module_config['base'], result, session)
return { return {
"result": result "result": result
} }
@@ -218,6 +232,8 @@ async def update(session: AsyncSessionDep, query_config, module_config, debug_da
create_convertor(module_config)['encode_list']([row]) create_convertor(module_config)['encode_list']([row])
await invoke_interceptor('before_update', module_config['base'], row, session)
try: try:
sql, values = build_update_sql(module_config, row, row.keys() if update_by_fields else None) sql, values = build_update_sql(module_config, row, row.keys() if update_by_fields else None)
@@ -232,6 +248,8 @@ async def update(session: AsyncSessionDep, query_config, module_config, debug_da
item_dict = get_value(result, 'result', None) item_dict = get_value(result, 'result', None)
await invoke_interceptor('after_update', module_config['base'], item_dict, session)
if item_dict is not None: if item_dict is not None:
return { return {
"result": item_dict "result": item_dict
@@ -268,6 +286,8 @@ async def batch_update(session: AsyncSessionDep, query_config, module_config, de
"rows": rows_without_id, "rows": rows_without_id,
} }
await invoke_interceptor('before_batch_update', module_config['base'], rows, session)
try: try:
for row in rows: for row in rows:
sql, values = build_update_sql(module_config, row, row.keys() if update_by_fields else None) sql, values = build_update_sql(module_config, row, row.keys() if update_by_fields else None)
@@ -283,6 +303,9 @@ async def batch_update(session: AsyncSessionDep, query_config, module_config, de
result = get_value(result, 'list', []) result = get_value(result, 'list', [])
if len(result) > 0: if len(result) > 0:
await invoke_interceptor('after_batch_update', module_config['base'], result, session)
return { return {
"result": result "result": result
} }
@@ -308,6 +331,8 @@ async def delete(session: AsyncSessionDep, query_config, module_config, debug_da
"error": "id parameter is missing", "error": "id parameter is missing",
} }
await invoke_interceptor('before_delete', module_config['base'], id, session)
try: try:
sql, values = build_delete_sql(module_config, id) sql, values = build_delete_sql(module_config, id)
debug_data.append({"sql": sql, "values": values}) debug_data.append({"sql": sql, "values": values})
@@ -316,6 +341,9 @@ async def delete(session: AsyncSessionDep, query_config, module_config, debug_da
deleted_rows = result.rowcount deleted_rows = result.rowcount
if deleted_rows >= 1: if deleted_rows >= 1:
await invoke_interceptor('after_delete', module_config['base'], id, session)
return {"deletedRows": deleted_rows} return {"deletedRows": deleted_rows}
else: else:
return {"error": f"delete failed, delete rows is {deleted_rows}", } return {"error": f"delete failed, delete rows is {deleted_rows}", }
@@ -0,0 +1,55 @@
from app.general.general_interceptors import add_general_interceptor, GeneralInterceptor
def add_general_interceptor_llm_user():
async def before_list(query_config, session):
print("before_list:llm_user", query_config, session)
async def after_list(rows, session):
print("after_list:llm_user", rows, session)
async def before_insert(row, session):
print("before_insert:llm_user", row, session)
async def after_insert(row, session):
print("after_insert:llm_user", row, session)
async def before_update(row, session):
print("before_update:llm_user", row, session)
async def after_update(row, session):
print("after_update:llm_user", row, session)
async def before_batch_insert(rows, session):
print("before_batch_insert:llm_user", rows, session)
async def after_batch_insert(rows, session):
print("after_batch_insert:llm_user", rows, session)
async def before_batch_update(rows, session):
print("before_batch_update:llm_user", rows, session)
async def after_batch_update(rows, session):
print("after_batch_update:llm_user", rows, session)
async def before_delete(id, session):
print("before_delete:llm_user", id, session)
async def after_delete(id, session):
print("after_delete:llm_user", id, session)
add_general_interceptor(GeneralInterceptor(
module="llm_user",
before_list=before_list,
after_list=after_list,
before_insert=before_insert,
after_insert=after_insert,
before_update=before_update,
after_update=after_update,
before_batch_insert=before_batch_insert,
after_batch_insert=after_batch_insert,
before_batch_update=before_batch_update,
after_batch_update=after_batch_update,
before_delete=before_delete,
after_delete=after_delete,
))
+139
View File
@@ -0,0 +1,139 @@
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_api_route import add_api_route
from app.controller.add_approve_route import add_approve_route
from app.controller.add_file_route import add_file_route
from app.controller.add_hotel_route import add_hotel_route
from app.controller.add_knowledge_route import add_knowledge_route
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_redis_route import add_redis_route
from app.controller.add_reimburse_route import add_reimburse_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.general.add_general_route import add_general_route
from app.general_interceptors.add_general_interceptor_llm_user import add_general_interceptor_llm_user
from app.model.ApiSecretModel import ApiSecretService
from app.model.ApproveModel import ApproveService
from app.model.ConversationModel import ConversationService
from app.model.HotelModel import HotelService
from app.model.InvoiceModel import InvoiceService
from app.model.KnowledgeBase import KnowledgeBaseService
from app.model.KnowledgeDoc import KnowledgeDocService, KnowledgeDocServiceWithCreator
from app.model.LgApprove import LgApproveService
from app.model.LgChat import LgChatService
from app.model.LgMessage import LgMessageService
from app.model.LlmOrder import LlmOrderService
from app.model.LlmProduct import LlmProductService
from app.model.ModuleModel import ModuleService
from app.model.OrderModel import OrderService
from app.model.OrgModel import OrgService
from app.model.PosModel import PosService
from app.model.ProjectModel import ProjectService
from app.model.ReimburseModel import ReimburseService
from app.model.ReimburseOtherModel import ReimburseOtherService
from app.model.ReimburseTravelModel import ReimburseTravelService
from app.model.RelProjUserModel import RelProjUserService
from app.utils.ModelInputSchema import ModelInputSchema
from app.utils.add_async_route import add_async_route
from app.utils.llm_utils import create_llm
from app.utils.next_id import add_next_id_route
app = create_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)
add_approve_route(app)
add_reimburse_route(app)
add_hotel_route(app)
add_file_route(app)
add_knowledge_route(app)
add_api_route(app)
add_redis_route(app)
add_general_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_routes(
app=app,
runnable=RunnableLambda(lambda x: x['messages']) | create_llm("doubao-vision-lite") | StrOutputParser(),
input_type=ModelInputSchema,
path="/doubao-vision-lite"
)
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")
OrgService.add_route(app=app, path="/org")
PosService.add_route(app=app, path="/pos")
ProjectService.add_route(app=app, path="/project")
RelProjUserService.add_route(app=app, path="/rel_proj_user")
ReimburseService.add_route(app=app, path="/reimburse")
ReimburseTravelService.add_route(app=app, path="/reimburse_travel")
ReimburseOtherService.add_route(app=app, path="/reimburse_other")
ApproveService.add_route(app=app, path="/approve")
HotelService.add_route(app=app, path="/hotel")
OrderService.add_route(app=app, path="/order")
InvoiceService.add_route(app=app, path="/invoice")
ConversationService.add_route(app=app, path="/conversation")
KnowledgeBaseService.add_route(app=app, path="/knowledge_base")
KnowledgeDocService.add_route(app=app, path="/knowledge_doc")
KnowledgeDocServiceWithCreator.add_route(app=app, path="/knowledge_doc_with_creator")
ApiSecretService.add_route(app=app, path="/api_secret")
ModuleService.add_route(app=app, path="/module")
add_general_interceptor_llm_user()
+10 -1
View File
@@ -1,8 +1,17 @@
import asyncio
import socket import socket
import sys
import psutil import psutil
from app.config.env import env from app.config.env import env
# 在Windows平台上设置事件循环策略为WindowsSelectorEventLoopPolicy
if sys.platform == "win32":
from asyncio import WindowsSelectorEventLoopPolicy
asyncio.set_event_loop_policy(WindowsSelectorEventLoopPolicy())
def run_uvicorn(): def run_uvicorn():
import uvicorn import uvicorn
@@ -25,7 +34,7 @@ uvicorn.run() 启动了一个异步事件循环来处理 HTTP 请求,
因此,uvicorn.run() 之后的代码不会被执行,直到服务器关闭。 因此,uvicorn.run() 之后的代码不会被执行,直到服务器关闭。
""" """
# 启动Uvicorn服务器 # 启动Uvicorn服务器
uvicorn.run("app.server:app", host="0.0.0.0", port=port) uvicorn.run("app.main_app:app", host="0.0.0.0", port=port)
# uvicorn.run之后的代码永远都不会执行 # uvicorn.run之后的代码永远都不会执行
# 使用FastAPI的 @app.on_event("startup")装饰器可以在服务器成功启动后执行代码 # 使用FastAPI的 @app.on_event("startup")装饰器可以在服务器成功启动后执行代码
+1 -136
View File
@@ -1,140 +1,5 @@
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_api_route import add_api_route
from app.controller.add_approve_route import add_approve_route
from app.controller.add_file_route import add_file_route
from app.controller.add_hotel_route import add_hotel_route
from app.controller.add_knowledge_route import add_knowledge_route
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_redis_route import add_redis_route
from app.controller.add_reimburse_route import add_reimburse_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.general.add_general_route import add_general_route
from app.model.ApiSecretModel import ApiSecretService
from app.model.ApproveModel import ApproveService
from app.model.ConversationModel import ConversationService
from app.model.HotelModel import HotelService
from app.model.InvoiceModel import InvoiceService
from app.model.KnowledgeBase import KnowledgeBaseService
from app.model.KnowledgeDoc import KnowledgeDocService, KnowledgeDocServiceWithCreator
from app.model.LgApprove import LgApproveService
from app.model.LgChat import LgChatService
from app.model.LgMessage import LgMessageService
from app.model.LlmOrder import LlmOrderService
from app.model.LlmProduct import LlmProductService
from app.model.ModuleModel import ModuleService
from app.model.OrderModel import OrderService
from app.model.OrgModel import OrgService
from app.model.PosModel import PosService
from app.model.ProjectModel import ProjectService
from app.model.ReimburseModel import ReimburseService
from app.model.ReimburseOtherModel import ReimburseOtherService
from app.model.ReimburseTravelModel import ReimburseTravelService
from app.model.RelProjUserModel import RelProjUserService
from app.run_uvicorn import run_uvicorn 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.llm_utils import create_llm
from app.utils.next_id import add_next_id_route
app = create_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)
add_approve_route(app)
add_reimburse_route(app)
add_hotel_route(app)
add_file_route(app)
add_knowledge_route(app)
add_api_route(app)
add_redis_route(app)
add_general_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_routes(
app=app,
runnable=RunnableLambda(lambda x: x['messages']) | create_llm("doubao-vision-lite") | StrOutputParser(),
input_type=ModelInputSchema,
path="/doubao-vision-lite"
)
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")
OrgService.add_route(app=app, path="/org")
PosService.add_route(app=app, path="/pos")
ProjectService.add_route(app=app, path="/project")
RelProjUserService.add_route(app=app, path="/rel_proj_user")
ReimburseService.add_route(app=app, path="/reimburse")
ReimburseTravelService.add_route(app=app, path="/reimburse_travel")
ReimburseOtherService.add_route(app=app, path="/reimburse_other")
ApproveService.add_route(app=app, path="/approve")
HotelService.add_route(app=app, path="/hotel")
OrderService.add_route(app=app, path="/order")
InvoiceService.add_route(app=app, path="/invoice")
ConversationService.add_route(app=app, path="/conversation")
KnowledgeBaseService.add_route(app=app, path="/knowledge_base")
KnowledgeDocService.add_route(app=app, path="/knowledge_doc")
KnowledgeDocServiceWithCreator.add_route(app=app, path="/knowledge_doc_with_creator")
ApiSecretService.add_route(app=app, path="/api_secret")
ModuleService.add_route(app=app, path="/module")
if __name__ == "__main__": if __name__ == "__main__":
run_uvicorn() run_uvicorn()
-7
View File
@@ -113,13 +113,6 @@ async def close_postgres_connection():
await PostgresCheckpointerManager.close_instance() 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): def create_test_graph(checkpointer: AsyncPostgresSaver):
class StateSchema(TypedDict): class StateSchema(TypedDict):
input: str input: str