From 44d64f82eaa605a37818892849a72befddc521eb Mon Sep 17 00:00:00 2001 From: wxg0103 <727495428@qq.com> Date: Tue, 22 Sep 2026 10:52:11 +0800 Subject: [PATCH] refactor: enhance model management with credential caching and retrieval methods --- apps/common/config/embedding_config.py | 49 ++++++++++++++++++++------ apps/models_provider/tests.py | 2 +- apps/models_provider/tools.py | 14 +++++--- 3 files changed, 49 insertions(+), 16 deletions(-) diff --git a/apps/common/config/embedding_config.py b/apps/common/config/embedding_config.py index 6d6bfde9593..cff306e737d 100644 --- a/apps/common/config/embedding_config.py +++ b/apps/common/config/embedding_config.py @@ -1,25 +1,49 @@ # coding=utf-8 """ - @project: maxkb - @Author:虎 - @file: embedding_config.py - @date:2023/10/23 16:03 - @desc: +@project: maxkb +@Author:虎 +@file: embedding_config.py +@date:2023/10/23 16:03 +@desc: """ import threading import time +import json from common.cache.mem_cache import MemCache +from common.utils.rsa_util import rsa_long_decrypt _lock = threading.Lock() locks = {} class ModelManage: - cache = MemCache('model', {}) + cache = MemCache("model", {}) + # 按 model_id 缓存的解密凭据与模型行,避免每次调用重复 RSA 解密 / 重复查库。 + # 二者在模型更新/删除时通过 delete_key(_id) 一并失效,保证改 key 立即生效。 + credential_cache = MemCache("model_credential", {}) + model_cache = MemCache("model_row", {}) up_clear_time = time.time() + @staticmethod + def get_decrypted_credential(model_id, credential): + """返回解密后的凭据 dict;结果按 model_id 缓存,改 key 时由 delete_key 清除。""" + cached = ModelManage.credential_cache.get(model_id) + if cached is not None: + return cached + decrypted = json.loads(rsa_long_decrypt(credential)) + ModelManage.credential_cache.set(model_id, decrypted, timeout=60 * 60 * 8) + return decrypted + + @staticmethod + def get_model_row(_id): + return ModelManage.model_cache.get(_id) + + @staticmethod + def set_model_row(_id, model): + ModelManage.model_cache.set(_id, model, timeout=60 * 60 * 8) + @staticmethod def _get_lock(_id): lock = locks.get(_id) @@ -59,24 +83,27 @@ def clear_timeout_cache(): @staticmethod def delete_key(_id): - if ModelManage.cache.has_key(_id): - ModelManage.cache.delete(_id) + for cache in (ModelManage.cache, ModelManage.credential_cache, ModelManage.model_cache): + if cache.has_key(_id): + cache.delete(_id) class VectorStore: from knowledge.vector.pg_vector import PGVector from knowledge.vector.base_vector import BaseVectorStore + instance_map = { - 'pg_vector': PGVector, + "pg_vector": PGVector, } instance = None @staticmethod def get_embedding_vector() -> BaseVectorStore: from knowledge.vector.pg_vector import PGVector + if VectorStore.instance is None: from maxkb.const import CONFIG - vector_store_class = VectorStore.instance_map.get(CONFIG.get("VECTOR_STORE_NAME"), - PGVector) + + vector_store_class = VectorStore.instance_map.get(CONFIG.get("VECTOR_STORE_NAME"), PGVector) VectorStore.instance = vector_store_class() return VectorStore.instance diff --git a/apps/models_provider/tests.py b/apps/models_provider/tests.py index b7f81022575..6f618128699 100644 --- a/apps/models_provider/tests.py +++ b/apps/models_provider/tests.py @@ -37,7 +37,7 @@ def test_every_registered_embedding_provider_declares_image_capability(self): embedding_classes = { model_info.model_class for provider in ModelProvideConstants - for model_info in provider.value.get_model_info_manage().model_list + for model_info in ModelProvideConstants[provider].get_model_info_manage().model_list if model_info.model_type == ModelTypeConst.EMBEDDING.name } diff --git a/apps/models_provider/tools.py b/apps/models_provider/tools.py index 58a6df7572d..01027d0177e 100644 --- a/apps/models_provider/tools.py +++ b/apps/models_provider/tools.py @@ -7,12 +7,10 @@ @desc: """ -import json from typing import Dict from common.config.embedding_config import ModelManage from common.database_model_manage.database_model_manage import DatabaseModelManage -from common.utils.rsa_util import rsa_long_decrypt from django.db import connection from django.db.models import QuerySet from django.utils.translation import gettext_lazy as _ @@ -35,7 +33,7 @@ def get_model_(provider, model_type, model_name, credential, model_id, use_local model = get_provider(provider).get_model( model_type, model_name, - json.loads(rsa_long_decrypt(credential)), + ModelManage.get_decrypted_credential(model_id, credential), model_id=model_id, use_local=use_local, streaming=True, @@ -110,14 +108,22 @@ def is_valid_credential( def get_model_by_id(_id, workspace_id): + # 同一工作空间读取模型行是高频路径,命中缓存可省一次 DB 查询。 + # 跨工作空间的授权读取不走缓存,始终保持走权限校验,避免越权风险。 + cached = ModelManage.get_model_row(_id) + if cached is not None and cached.workspace_id == workspace_id: + return cached model = QuerySet(Model).filter(id=_id).first() # 归还链接到连接池 connection.close() get_authorized_model = DatabaseModelManage.get_model("get_authorized_model") - if model and model.workspace_id != workspace_id and get_authorized_model is not None: + authorized = model is not None and model.workspace_id != workspace_id and get_authorized_model is not None + if authorized: model = get_authorized_model(QuerySet(Model).filter(id=_id), workspace_id).first() if model is None: raise Exception(_("Model does not exist")) + if not authorized: + ModelManage.set_model_row(_id, model) return model