feat: 秘钥凭证管理
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user