diff --git a/app/controller/add_api_route.py b/app/controller/add_api_route.py new file mode 100644 index 0000000..08e0686 --- /dev/null +++ b/app/controller/add_api_route.py @@ -0,0 +1,12 @@ +from fastapi import FastAPI +from starlette.requests import Request + + +def add_api_route(app: FastAPI): + # 用于dify自定义插件授权认证接口 + @app.post('/api/account') + async def api_account( + request: Request + ): + print("request.state", request.state) + return request.state.user diff --git a/app/middlewares/app_middlewares.py b/app/middlewares/app_middlewares.py index 6229f8e..dfe3b75 100644 --- a/app/middlewares/app_middlewares.py +++ b/app/middlewares/app_middlewares.py @@ -9,6 +9,8 @@ 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.model.ApiSecretModel import ApiSecretService +from app.utils.api_secret_utils import api_secret_utils, ApiSecretStatus from app.utils.db_utils import async_session @@ -58,6 +60,30 @@ def add_app_middlewares(app: FastAPI): except InvalidTokenError: raise unauthorized_exception + # 如果请求的是api接口 + # api开头的接口需要额外验证秘钥 + if request.url.path.startswith("/api/"): + # 验证秘钥状态 + secret_status = await api_secret_utils.verify_secret(token) + print("secret_status", secret_status) + if secret_status == ApiSecretStatus.invalid: + # 秘钥无效 + raise unauthorized_exception + elif secret_status == ApiSecretStatus.not_exist: + # 秘钥没有缓存 + async with async_session() as session: + query_cls = await ApiSecretService.query_item(session, {"secret": token}) + if query_cls: + # 秘钥存在 + await api_secret_utils.save_secret(token, ApiSecretStatus.valid) + else: + # 秘钥不存在 + await api_secret_utils.save_secret(token, ApiSecretStatus.invalid) + raise unauthorized_exception + else: + # 秘钥有效 + pass + response = await call_next(request) return response diff --git a/app/model/ApiSecretModel.py b/app/model/ApiSecretModel.py new file mode 100644 index 0000000..2677860 --- /dev/null +++ b/app/model/ApiSecretModel.py @@ -0,0 +1,52 @@ +import uuid +from datetime import timedelta + +from sqlmodel import Field + +from app.model.BasicModel import BasicModel +from app.model.UserModel import UserService, UserServiceModel +from app.utils.CrpyUtils import CryptUtils +from app.utils.api_secret_utils import api_secret_utils +from app.utils.create_module_service import create_model_service + + +class ApiSecretModel(BasicModel, table=True): + __tablename__ = "pl_api_secret" + + secret: str = Field(..., description="api秘钥") + description: str = Field(..., description="描述信息") + + +# 插入数据之前,查询当前用户信息,使用username生成一个access_token作为秘钥保存到 ApiSecretModel中 +async def before_insert(row_dict, session): + user: UserServiceModel = await UserService.query_item(session, {"id": row_dict.get("created_by")}) + row_dict["secret"] = CryptUtils.create_token(user.username, "access", timedelta(days=9999)) + + +# 删除ApiSecret凭据之前,先清理掉缓存中的秘钥信息 +async def before_delete(row_dict, session): + row_cls: ApiSecretModel = await ApiSecretService.query_item(session, {"id": row_dict.get("id")}) + await api_secret_utils.remove_secret(row_cls.secret) + + +# 更新ApiSecret凭据之前,先清理掉缓存中的秘钥信息 +async def before_update(row_dict, session): + row_cls: ApiSecretModel = await ApiSecretService.query_item(session, {"id": row_dict.get("id")}) + await api_secret_utils.remove_secret(row_cls.secret) + + +# 不允许批量处理ApiSecret凭据信息 +async def before_batch_method(row_dict_list, session): + raise Exception("秘钥不支持批量操作!") + + +ApiSecretService = create_model_service( + Cls=ApiSecretModel, + before_insert=before_insert, + before_update=before_update, + before_delete=before_delete, + + before_batch_insert=before_batch_method, + before_batch_update=before_batch_method, + before_batch_delete=before_batch_method, +) diff --git a/app/server.py b/app/server.py index 06ae1bc..06f5919 100644 --- a/app/server.py +++ b/app/server.py @@ -5,6 +5,7 @@ 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 @@ -23,6 +24,7 @@ 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.model.ApiSecretModel import ApiSecretService from app.model.ApproveModel import ApproveService from app.model.ConversationModel import ConversationService from app.model.HotelModel import HotelService @@ -68,6 +70,7 @@ add_reimburse_route(app) add_hotel_route(app) add_file_route(app) add_knowledge_route(app) +add_api_route(app) @app.get("/get_env") @@ -125,6 +128,7 @@ 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") if __name__ == "__main__": run_uvicorn() diff --git a/app/utils/api_secret_utils.py b/app/utils/api_secret_utils.py new file mode 100644 index 0000000..105bda8 --- /dev/null +++ b/app/utils/api_secret_utils.py @@ -0,0 +1,34 @@ +from enum import Enum + + +class ApiSecretStatus(Enum): + valid = "valid" # 秘钥有效 + invalid = "invalid" # 秘钥无效 + not_exist = "not_exist" # 秘钥不存在 + + +class ApiSecretUtils(): + def __init__(self): + self.cache = {} + + # 验证秘钥状态 + async def verify_secret(self, secret: str) -> ApiSecretStatus: + api_secret_status = self.cache.get(secret, None) + print("verify_secret", self.cache) + return api_secret_status or ApiSecretStatus.not_exist + + # 保存秘钥状态 + async def save_secret(self, secret: str, status: ApiSecretStatus): + print("save_secret", self.cache) + self.cache[secret] = status + + # 移除秘钥缓存 + async def remove_secret(self, secret: str): + # self.cache.pop(secret, None) + if secret in self.cache: + del self.cache[secret] + print("remove_secret", secret) + print(await self.verify_secret(secret)) + + +api_secret_utils = ApiSecretUtils() diff --git a/app/utils/create_module_service.py b/app/utils/create_module_service.py index d71b152..c656059 100644 --- a/app/utils/create_module_service.py +++ b/app/utils/create_module_service.py @@ -309,11 +309,11 @@ def create_model_service( row_dict: dict = Body(..., description=f"插入的数据,字段参考{Cls.__name__}"), user: ModelPublicUser = None, ): + await self.check_insert_row_dict(row_dict, user) + if before_insert is not None: await before_insert(row_dict, session) - await self.check_insert_row_dict(row_dict, user) - try: # 使用模型类验证数据并创建实例(校验字段类型和约束) insert_cls = Cls.model_validate(row_dict) @@ -342,9 +342,6 @@ def create_model_service( user: ModelPublicUser = None, ): - if before_batch_insert is not None: - await before_batch_insert(row_dict_list, session) - # 筛选出没有id的记录(需要自动生成id) row_dict_list_without_id = [] @@ -363,6 +360,9 @@ def create_model_service( for row_dict in row_dict_list_without_id: await self.check_insert_row_dict(row_dict, user) + if before_batch_insert is not None: + await before_batch_insert(row_dict_list, session) + try: # 验证所有记录并转换为模型实例列表 insert_cls_list = [Cls.model_validate(row_dict) for row_dict in row_dict_list]