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
117 changes: 68 additions & 49 deletions apps/models_provider/constants/model_provider_constants.py
Original file line number Diff line number Diff line change
@@ -1,52 +1,71 @@
# coding=utf-8
from enum import Enum
import importlib
import threading
from typing import Iterator

from models_provider.impl.aliyun_bai_lian_model_provider.aliyun_bai_lian_model_provider import (
AliyunBaiLianModelProvider,
)
from models_provider.impl.anthropic_model_provider.anthropic_model_provider import AnthropicModelProvider
from models_provider.impl.aws_bedrock_model_provider.aws_bedrock_model_provider import BedrockModelProvider
from models_provider.impl.azure_model_provider.azure_model_provider import AzureModelProvider
from models_provider.impl.deepseek_model_provider.deepseek_model_provider import DeepSeekModelProvider
from models_provider.impl.docker_ai_model_provider.docker_ai_model_provider import DockerModelProvider
from models_provider.impl.gemini_model_provider.gemini_model_provider import GeminiModelProvider
from models_provider.impl.kimi_model_provider.kimi_model_provider import KimiModelProvider
from models_provider.impl.local_model_provider.local_model_provider import LocalModelProvider
from models_provider.impl.minimax_model_provider.minimax_model_provider import MiniMaxModelProvider
from models_provider.impl.ollama_model_provider.ollama_model_provider import OllamaModelProvider
from models_provider.impl.openai_model_provider.openai_model_provider import OpenAIModelProvider
from models_provider.impl.regolo_model_provider.regolo_model_provider import RegoloModelProvider
from models_provider.impl.siliconCloud_model_provider.siliconCloud_model_provider import SiliconCloudModelProvider
from models_provider.impl.tencent_model_provider.tencent_model_provider import TencentModelProvider
from models_provider.impl.vllm_model_provider.vllm_model_provider import VllmModelProvider
from models_provider.impl.volcanic_engine_model_provider.volcanic_engine_model_provider import (
VolcanicEngineModelProvider,
# 供应商注册表:模型供应商不再用 Enum 硬编码实例,而是注册为"模块路径 + 类名"的
# 惰性工厂。__getitem__ 首次访问某供应商时才 importlib 加载对应模块并缓存实例,
# 因此导入 models_provider 时不会连带加载 21 个供应商的重依赖(openai/bedrock/
# gemini/azure 等),只在真正使用某个供应商时按需加载,从而加快启动、降低内存。


class _ProviderRegistry:
def __init__(self):
self._factories: dict[str, tuple[str, str]] = {}
self._instances: dict[str, object] = {}
self._lock = threading.Lock()

def register(self, name: str, provider_dir: str, class_name: str) -> "_ProviderRegistry":
self._factories[name] = (f"models_provider.impl.{provider_dir}.{provider_dir}", class_name)
return self

def __getitem__(self, name: str):
instance = self._instances.get(name)
if instance is None:
with self._lock:
instance = self._instances.get(name)
if instance is None:
module_path, class_name = self._factories[name]
module = importlib.import_module(module_path)
instance = getattr(module, class_name)()
self._instances[name] = instance
return instance

def __iter__(self) -> Iterator[str]:
return iter(self._factories)

def __contains__(self, name: object) -> bool:
return name in self._factories

def __len__(self) -> int:
return len(self._factories)

@property
def __members__(self) -> dict:
return self._factories


ModelProvideConstants = (
_ProviderRegistry()
.register("model_azure_provider", "azure_model_provider", "AzureModelProvider")
.register("model_qianfan_provider", "qianfan_model_provider", "QianfanModelProvider")
.register("model_ollama_provider", "ollama_model_provider", "OllamaModelProvider")
.register("model_openai_provider", "openai_model_provider", "OpenAIModelProvider")
.register("model_docker_ai_provider", "docker_ai_model_provider", "DockerModelProvider")
.register("model_kimi_provider", "kimi_model_provider", "KimiModelProvider")
.register("model_zhipu_provider", "zhipu_model_provider", "ZhiPuModelProvider")
.register("model_xf_provider", "xf_model_provider", "XunFeiModelProvider")
.register("model_deepseek_provider", "deepseek_model_provider", "DeepSeekModelProvider")
.register("model_gemini_provider", "gemini_model_provider", "GeminiModelProvider")
.register("model_volcanic_engine_provider", "volcanic_engine_model_provider", "VolcanicEngineModelProvider")
.register("model_tencent_provider", "tencent_model_provider", "TencentModelProvider")
.register("model_aws_bedrock_provider", "aws_bedrock_model_provider", "BedrockModelProvider")
.register("model_local_provider", "local_model_provider", "LocalModelProvider")
.register("model_xinference_provider", "xinference_model_provider", "XinferenceModelProvider")
.register("model_vllm_provider", "vllm_model_provider", "VllmModelProvider")
.register("aliyun_bai_lian_model_provider", "aliyun_bai_lian_model_provider", "AliyunBaiLianModelProvider")
.register("model_anthropic_provider", "anthropic_model_provider", "AnthropicModelProvider")
.register("model_siliconCloud_provider", "siliconCloud_model_provider", "SiliconCloudModelProvider")
.register("model_regolo_provider", "regolo_model_provider", "RegoloModelProvider")
.register("model_minimax_provider", "minimax_model_provider", "MiniMaxModelProvider")
)
from models_provider.impl.qianfan_model_provider.qianfan_model_provider import QianfanModelProvider
from models_provider.impl.xf_model_provider.xf_model_provider import XunFeiModelProvider
from models_provider.impl.xinference_model_provider.xinference_model_provider import XinferenceModelProvider
from models_provider.impl.zhipu_model_provider.zhipu_model_provider import ZhiPuModelProvider


class ModelProvideConstants(Enum):
model_azure_provider = AzureModelProvider()
model_qianfan_provider = QianfanModelProvider()
model_ollama_provider = OllamaModelProvider()
model_openai_provider = OpenAIModelProvider()
model_docker_ai_provider = DockerModelProvider()
model_kimi_provider = KimiModelProvider()
model_zhipu_provider = ZhiPuModelProvider()
model_xf_provider = XunFeiModelProvider()
model_deepseek_provider = DeepSeekModelProvider()
model_gemini_provider = GeminiModelProvider()
model_volcanic_engine_provider = VolcanicEngineModelProvider()
model_tencent_provider = TencentModelProvider()
model_aws_bedrock_provider = BedrockModelProvider()
model_local_provider = LocalModelProvider()
model_xinference_provider = XinferenceModelProvider()
model_vllm_provider = VllmModelProvider()
aliyun_bai_lian_model_provider = AliyunBaiLianModelProvider()
model_anthropic_provider = AnthropicModelProvider()
model_siliconCloud_provider = SiliconCloudModelProvider()
model_regolo_provider = RegoloModelProvider()
model_minimax_provider = MiniMaxModelProvider()
27 changes: 15 additions & 12 deletions apps/models_provider/impl/base_chat_open_ai.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,10 @@ def custom_get_token_ids(text: str):
return tokenizer.encode(text)


# 复用固定线程池,避免每次 token 计数都新建/销毁线程;线程只在首次调用时创建
_token_count_executor = ThreadPoolExecutor(max_workers=4, thread_name_prefix="maxkb-token-count")


def _convert_delta_to_message_chunk(
_dict: Mapping[str, Any], default_class: type[BaseMessageChunk]
) -> BaseMessageChunk:
Expand Down Expand Up @@ -100,18 +104,17 @@ def get_num_tokens_from_messages(
timeout: Optional[float] = 0.5,
) -> int:
if self.usage_metadata is None or self.usage_metadata == {}:
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(super().get_num_tokens_from_messages, messages, tools)
try:
response = future.result(timeout=timeout)
maxkb_logger.info("请求成功(未超时)")
return response
except Exception as e:
if isinstance(e, ReadTimeout):
raise # 继续抛出
else:
tokenizer = TokenizerManage.get_tokenizer()
return sum([len(tokenizer.encode(get_buffer_string([m]))) for m in messages])
future = _token_count_executor.submit(super().get_num_tokens_from_messages, messages, tools)
try:
response = future.result(timeout=timeout)
maxkb_logger.info("请求成功(未超时)")
return response
except Exception as e:
if isinstance(e, ReadTimeout):
raise # 继续抛出
else:
tokenizer = TokenizerManage.get_tokenizer()
return sum([len(tokenizer.encode(get_buffer_string([m]))) for m in messages])

return self.usage_metadata.get("input_tokens", self.usage_metadata.get("prompt_tokens", 0))

Expand Down
12 changes: 5 additions & 7 deletions apps/models_provider/serializers/model_serializer.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,9 +72,7 @@ class ModelPullManage:
@staticmethod
def pull(model: Model, credential: Dict):
try:
response = ModelProvideConstants[model.provider].value.down_model(
model.model_type, model.model_name, credential
)
response = ModelProvideConstants[model.provider].down_model(model.model_type, model.model_name, credential)
down_model_chunk = {}
last_update_time = time.time()

Expand Down Expand Up @@ -117,7 +115,7 @@ def model_to_dict(model: Model):
"status": model.status,
"meta": model.meta,
"credential": ModelProvideConstants[model.provider]
.value.get_model_credential(model.model_type, model.model_name)
.get_model_credential(model.model_type, model.model_name)
.encryption_dict(credential),
"workspace_id": model.workspace_id,
"nick_name": model.user.nick_name if model.user else "",
Expand Down Expand Up @@ -272,8 +270,8 @@ def is_valid(self, model=None, raise_exception=False):
model_type = self.data.get("model_type")
model_name = self.data.get("model_name")
credential = self.data.get("credential")
provider_handler = ModelProvideConstants[provider].value
model_credential = ModelProvideConstants[provider].value.get_model_credential(model_type, model_name)
provider_handler = ModelProvideConstants[provider]
model_credential = ModelProvideConstants[provider].get_model_credential(model_type, model_name)
source_model_credential = json.loads(rsa_long_decrypt(model.credential))
source_encryption_model_credential = model_credential.encryption_dict(source_model_credential)
if credential is not None:
Expand Down Expand Up @@ -303,7 +301,7 @@ def is_valid(self, *, raise_exception=False):
500, _("base model【{model_name}】already exists").format(model_name=self.data.get("name"))
)
default_params = {item["field"]: item["default_value"] for item in self.data.get("model_params_form")}
ModelProvideConstants[self.data.get("provider")].value.is_valid_credential(
ModelProvideConstants[self.data.get("provider")].is_valid_credential(
self.data.get("model_type"),
self.data.get("model_name"),
self.data.get("credential"),
Expand Down
2 changes: 1 addition & 1 deletion apps/models_provider/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,7 +59,7 @@ def get_provider(provider):
@param provider: 供应商字符串
@return: 供应商实例
"""
return ModelProvideConstants[provider].value
return ModelProvideConstants[provider]


def get_model_list(provider, model_type):
Expand Down
15 changes: 6 additions & 9 deletions apps/models_provider/views/provide.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,19 +33,16 @@ def get(self, request: Request):
len(
[
item
for item in ModelProvideConstants[key].value.get_model_type_list()
for item in ModelProvideConstants[key].get_model_type_list()
if item["value"] == model_type
]
)
> 0
):
providers.append(ModelProvideConstants[key].value.get_model_provide_info().to_dict())
providers.append(ModelProvideConstants[key].get_model_provide_info().to_dict())
return result.success(providers)
return result.success(
[
ModelProvideConstants[key].value.get_model_provide_info().to_dict()
for key in ModelProvideConstants.__members__
]
[ModelProvideConstants[key].get_model_provide_info().to_dict() for key in ModelProvideConstants.__members__]
)

class ModelTypeList(APIView):
Expand All @@ -62,7 +59,7 @@ class ModelTypeList(APIView):
)
def get(self, request: Request):
provider = request.query_params.get("provider")
return result.success(ModelProvideConstants[provider].value.get_model_type_list())
return result.success(ModelProvideConstants[provider].get_model_type_list())

class ModelList(APIView):
authentication_classes = [TokenAuth]
Expand All @@ -80,7 +77,7 @@ def get(self, request: Request):
provider = request.query_params.get("provider")
model_type = request.query_params.get("model_type")

return result.success(ModelProvideConstants[provider].value.get_model_list(model_type))
return result.success(ModelProvideConstants[provider].get_model_list(model_type))

class ModelParamsForm(APIView):
authentication_classes = [TokenAuth]
Expand Down Expand Up @@ -118,5 +115,5 @@ def get(self, request: Request):
model_type = request.query_params.get("model_type")
model_name = request.query_params.get("model_name")
return result.success(
ModelProvideConstants[provider].value.get_model_credential(model_type, model_name).to_form_list()
ModelProvideConstants[provider].get_model_credential(model_type, model_name).to_form_list()
)
Loading