Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 38 additions & 11 deletions apps/common/config/embedding_config.py
Original file line number Diff line number Diff line change
@@ -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)
Expand Down Expand Up @@ -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
2 changes: 1 addition & 1 deletion apps/models_provider/tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
}

Expand Down
14 changes: 10 additions & 4 deletions apps/models_provider/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 _
Expand All @@ -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,
Expand Down Expand Up @@ -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


Expand Down
Loading