2024-07-07 01:58:23 +08:00
|
|
|
|
import os
|
|
|
|
|
|
import json
|
|
|
|
|
|
import yaml
|
2024-11-14 20:07:49 +08:00
|
|
|
|
from pathlib import Path
|
2025-02-27 19:35:25 +08:00
|
|
|
|
from src.utils.logging_config import logger
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
2025-02-28 00:26:57 +08:00
|
|
|
|
with open(Path("src/static/models.yaml"), 'r', encoding='utf-8') as f:
|
2024-11-14 20:07:49 +08:00
|
|
|
|
_models = yaml.safe_load(f)
|
|
|
|
|
|
|
|
|
|
|
|
MODEL_NAMES = _models["MODEL_NAMES"]
|
|
|
|
|
|
EMBED_MODEL_INFO = _models["EMBED_MODEL_INFO"]
|
|
|
|
|
|
RERANKER_LIST = _models["RERANKER_LIST"]
|
|
|
|
|
|
|
2025-03-14 03:28:53 +08:00
|
|
|
|
DEFAULT_MOCK_API = 'this_is_mock_api_key_in_frontend'
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
class SimpleConfig(dict):
|
|
|
|
|
|
|
|
|
|
|
|
def __key(self, key):
|
2025-03-29 17:33:09 +08:00
|
|
|
|
return "" if key is None else key # 目前忘记了这里为什么要 lower 了,只能说配置项最好不要有大写的
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
def __str__(self):
|
|
|
|
|
|
return json.dumps(self)
|
|
|
|
|
|
|
|
|
|
|
|
def __setattr__(self, key, value):
|
|
|
|
|
|
self[self.__key(key)] = value
|
|
|
|
|
|
|
|
|
|
|
|
def __getattr__(self, key):
|
|
|
|
|
|
return self.get(self.__key(key))
|
|
|
|
|
|
|
|
|
|
|
|
def __getitem__(self, key):
|
|
|
|
|
|
return super().get(self.__key(key))
|
|
|
|
|
|
|
|
|
|
|
|
def __setitem__(self, key, value):
|
|
|
|
|
|
return super().__setitem__(self.__key(key), value)
|
|
|
|
|
|
|
2024-07-25 20:30:28 +08:00
|
|
|
|
def __dict__(self):
|
|
|
|
|
|
return {k: v for k, v in self.items()}
|
|
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
class Config(SimpleConfig):
|
|
|
|
|
|
|
2025-02-20 01:26:12 +08:00
|
|
|
|
def __init__(self):
|
2024-07-07 01:58:23 +08:00
|
|
|
|
super().__init__()
|
2024-07-25 20:30:28 +08:00
|
|
|
|
self._config_items = {}
|
2025-02-20 01:26:12 +08:00
|
|
|
|
self.save_dir = "saves"
|
2025-02-26 23:58:26 +08:00
|
|
|
|
self.filename = str(Path("saves/config/base.yaml"))
|
2025-02-20 01:26:12 +08:00
|
|
|
|
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
2024-07-17 18:52:20 +08:00
|
|
|
|
### >>> 默认配置
|
2024-07-25 20:30:28 +08:00
|
|
|
|
self.add_item("stream", default=True, des="是否开启流式输出")
|
2024-07-17 18:52:20 +08:00
|
|
|
|
# 功能选项
|
2024-07-29 01:00:02 +08:00
|
|
|
|
self.add_item("enable_reranker", default=False, des="是否开启重排序")
|
2024-07-31 20:22:05 +08:00
|
|
|
|
self.add_item("enable_knowledge_base", default=False, des="是否开启知识库")
|
2024-09-03 16:37:59 +08:00
|
|
|
|
self.add_item("enable_knowledge_graph", default=False, des="是否开启知识图谱")
|
2025-02-24 14:58:35 +08:00
|
|
|
|
self.add_item("enable_web_search", default=False, des="是否开启网页搜索(需配置 TAVILY_API_KEY)")
|
2024-07-17 18:52:20 +08:00
|
|
|
|
# 模型配置
|
|
|
|
|
|
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
|
2025-02-20 01:26:12 +08:00
|
|
|
|
## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR
|
2025-02-23 16:38:56 +08:00
|
|
|
|
self.add_item("model_provider", default="siliconflow", des="模型提供商", choices=list(MODEL_NAMES.keys()))
|
2025-02-25 21:26:37 +08:00
|
|
|
|
self.add_item("model_provider_lite", default="siliconflow", des="模型提供商(用于轻量任务)", choices=list(MODEL_NAMES.keys()))
|
2025-02-23 16:38:56 +08:00
|
|
|
|
self.add_item("model_name", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称")
|
2025-02-25 21:26:37 +08:00
|
|
|
|
self.add_item("model_name_lite", default="Qwen/Qwen2.5-7B-Instruct", des="模型名称(用于轻量任务)")
|
2025-02-26 23:58:26 +08:00
|
|
|
|
|
2025-02-23 16:38:56 +08:00
|
|
|
|
self.add_item("embed_model", default="siliconflow/BAAI/bge-m3", des="Embedding 模型", choices=list(EMBED_MODEL_INFO.keys()))
|
|
|
|
|
|
self.add_item("reranker", default="siliconflow/BAAI/bge-reranker-v2-m3", des="Re-Ranker 模型", choices=list(RERANKER_LIST.keys()))
|
2024-07-25 20:30:28 +08:00
|
|
|
|
self.add_item("model_local_paths", default={}, des="本地模型路径")
|
2025-02-20 01:26:12 +08:00
|
|
|
|
self.add_item("use_rewrite_query", default="off", des="重写查询", choices=["off", "on", "hyde"])
|
2025-03-07 01:05:50 +08:00
|
|
|
|
self.add_item("device", default="cuda", des="运行本地模型的设备", choices=["cpu", "cuda"])
|
2024-07-17 18:52:20 +08:00
|
|
|
|
### <<< 默认配置结束
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
self.load()
|
2024-07-09 05:04:20 +08:00
|
|
|
|
self.handle_self()
|
|
|
|
|
|
|
2024-07-25 20:30:28 +08:00
|
|
|
|
def add_item(self, key, default, des=None, choices=None):
|
|
|
|
|
|
self.__setattr__(key, default)
|
|
|
|
|
|
self._config_items[key] = {
|
|
|
|
|
|
"default": default,
|
|
|
|
|
|
"des": des,
|
|
|
|
|
|
"choices": choices
|
|
|
|
|
|
}
|
|
|
|
|
|
|
2024-09-03 16:37:59 +08:00
|
|
|
|
def __dict__(self):
|
|
|
|
|
|
blocklist = [
|
|
|
|
|
|
"_config_items",
|
|
|
|
|
|
"model_names",
|
2024-09-28 00:41:18 +08:00
|
|
|
|
"model_provider_status",
|
2025-03-04 13:49:00 +08:00
|
|
|
|
"embed_model_names",
|
|
|
|
|
|
"reranker_names",
|
2024-09-03 16:37:59 +08:00
|
|
|
|
]
|
|
|
|
|
|
return {k: v for k, v in self.items() if k not in blocklist}
|
|
|
|
|
|
|
2024-07-09 05:04:20 +08:00
|
|
|
|
def handle_self(self):
|
2024-07-31 20:22:05 +08:00
|
|
|
|
self.model_names = MODEL_NAMES
|
2025-02-28 02:39:47 +08:00
|
|
|
|
self.embed_model_names = EMBED_MODEL_INFO
|
|
|
|
|
|
self.reranker_names = RERANKER_LIST
|
|
|
|
|
|
|
2024-10-08 22:16:17 +08:00
|
|
|
|
model_provider_info = self.model_names.get(self.model_provider, {})
|
2025-02-20 01:26:12 +08:00
|
|
|
|
self.model_dir = os.environ.get("MODEL_DIR", "")
|
2025-03-11 14:18:48 +08:00
|
|
|
|
logger.info(f"MODEL_DIR: {self.model_dir}; 如果是在 docker 中运行,会自动挂载 MODEL_DIR 到 /models 目录,请检查 docker compose 文件")
|
2024-07-31 20:22:05 +08:00
|
|
|
|
|
2025-02-20 01:26:12 +08:00
|
|
|
|
# 检查模型提供商是否存在
|
2024-10-08 22:16:17 +08:00
|
|
|
|
if self.model_provider != "custom":
|
|
|
|
|
|
if self.model_name not in model_provider_info["models"]:
|
|
|
|
|
|
logger.warning(f"Model name {self.model_name} not in {self.model_provider}, using default model name")
|
|
|
|
|
|
self.model_name = model_provider_info["default"]
|
2024-07-31 20:22:05 +08:00
|
|
|
|
|
2024-10-08 22:16:17 +08:00
|
|
|
|
default_model_name = model_provider_info["default"]
|
|
|
|
|
|
self.model_name = self.get("model_name") or default_model_name
|
|
|
|
|
|
else:
|
|
|
|
|
|
self.model_name = self.get("model_name")
|
2025-03-29 17:46:19 +08:00
|
|
|
|
if self.model_name not in [item["custom_id"] for item in self.get("custom_models", [])]:
|
2024-10-08 22:16:17 +08:00
|
|
|
|
logger.warning(f"Model name {self.model_name} not in custom models, using default model name")
|
2025-03-29 17:46:19 +08:00
|
|
|
|
if self.get("custom_models", []):
|
|
|
|
|
|
self.model_name = self.get("custom_models", [])[0]["custom_id"]
|
|
|
|
|
|
else:
|
|
|
|
|
|
self.model_name = self._config_items["model_name"]["default"]
|
|
|
|
|
|
self.model_provider = self._config_items["model_provider"]["default"]
|
|
|
|
|
|
logger.error(f"No custom models found, using default model {self.model_name} from {self.model_provider}")
|
2024-07-31 20:22:05 +08:00
|
|
|
|
|
2025-02-20 01:26:12 +08:00
|
|
|
|
# 检查模型提供商的环境变量
|
2024-11-14 20:07:49 +08:00
|
|
|
|
conds = {}
|
2024-09-28 00:41:18 +08:00
|
|
|
|
self.model_provider_status = {}
|
|
|
|
|
|
for provider in self.model_names:
|
2024-11-14 20:07:49 +08:00
|
|
|
|
conds[provider] = self.model_names[provider]["env"]
|
|
|
|
|
|
conds_bool = [bool(os.getenv(_k)) for _k in conds[provider]]
|
|
|
|
|
|
self.model_provider_status[provider] = all(conds_bool)
|
|
|
|
|
|
|
2025-02-20 01:26:12 +08:00
|
|
|
|
# 检查web_search的环境变量
|
|
|
|
|
|
if self.enable_web_search and not os.getenv("TAVILY_API_KEY"):
|
|
|
|
|
|
logger.warning("TAVILY_API_KEY not set, web search will be disabled")
|
|
|
|
|
|
self.enable_web_search = False
|
|
|
|
|
|
|
2024-11-14 20:07:49 +08:00
|
|
|
|
self.valuable_model_provider = [k for k, v in self.model_provider_status.items() if v]
|
|
|
|
|
|
assert len(self.valuable_model_provider) > 0, f"No model provider available, please check your `.env` file. API_KEY_LIST: {conds}"
|
|
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
def load(self):
|
2024-07-17 18:52:20 +08:00
|
|
|
|
"""根据传入的文件覆盖掉默认配置"""
|
2024-07-22 00:00:54 +08:00
|
|
|
|
logger.info(f"Loading config from {self.filename}")
|
2024-07-07 01:58:23 +08:00
|
|
|
|
if self.filename is not None and os.path.exists(self.filename):
|
2024-08-25 20:29:24 +08:00
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
if self.filename.endswith(".json"):
|
|
|
|
|
|
with open(self.filename, 'r') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
content = f.read()
|
|
|
|
|
|
if content:
|
2024-08-25 20:29:24 +08:00
|
|
|
|
local_config = json.loads(content)
|
|
|
|
|
|
self.update(local_config)
|
2024-07-25 20:30:28 +08:00
|
|
|
|
else:
|
|
|
|
|
|
print(f"{self.filename} is empty.")
|
2024-08-25 20:29:24 +08:00
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
elif self.filename.endswith(".yaml"):
|
|
|
|
|
|
with open(self.filename, 'r') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
content = f.read()
|
|
|
|
|
|
if content:
|
2024-08-25 20:29:24 +08:00
|
|
|
|
local_config = yaml.safe_load(content)
|
|
|
|
|
|
self.update(local_config)
|
2024-07-25 20:30:28 +08:00
|
|
|
|
else:
|
|
|
|
|
|
print(f"{self.filename} is empty.")
|
|
|
|
|
|
else:
|
|
|
|
|
|
logger.warning(f"Unknown config file type {self.filename}")
|
2024-08-25 20:29:24 +08:00
|
|
|
|
|
2024-07-07 01:58:23 +08:00
|
|
|
|
else:
|
2024-07-28 16:16:52 +08:00
|
|
|
|
logger.warning(f"\n\n{'='*70}\n{'Config file not found':^70}\n{'You can custum your config in `' + self.filename + '`':^70}\n{'='*70}\n\n")
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
|
def save(self):
|
2024-07-22 00:00:54 +08:00
|
|
|
|
logger.info(f"Saving config to {self.filename}")
|
|
|
|
|
|
if self.filename is None:
|
|
|
|
|
|
logger.warning("Config file is not specified, save to default config/base.yaml")
|
2025-02-26 23:58:26 +08:00
|
|
|
|
self.filename = os.path.join(self.save_dir, "config", "base.yaml")
|
2024-07-29 01:00:02 +08:00
|
|
|
|
os.makedirs(os.path.dirname(self.filename), exist_ok=True)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
|
|
|
|
|
if self.filename.endswith(".json"):
|
|
|
|
|
|
with open(self.filename, 'w+') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
json.dump(self.__dict__(), f, indent=4, ensure_ascii=False)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
elif self.filename.endswith(".yaml"):
|
|
|
|
|
|
with open(self.filename, 'w+') as f:
|
2024-07-25 20:30:28 +08:00
|
|
|
|
yaml.dump(self.__dict__(), f, indent=2, allow_unicode=True)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
else:
|
|
|
|
|
|
logger.warning(f"Unknown config file type {self.filename}, save as json")
|
|
|
|
|
|
with open(self.filename, 'w+') as f:
|
2024-07-07 01:58:23 +08:00
|
|
|
|
json.dump(self, f, indent=4)
|
2024-07-22 00:00:54 +08:00
|
|
|
|
|
2024-07-31 20:22:05 +08:00
|
|
|
|
logger.info(f"Config file {self.filename} saved")
|
2025-03-14 03:28:53 +08:00
|
|
|
|
|
|
|
|
|
|
def get_safe_config(self):
|
|
|
|
|
|
"""
|
|
|
|
|
|
获取安全的配置,即过滤掉 api_key
|
|
|
|
|
|
"""
|
2025-03-29 17:33:09 +08:00
|
|
|
|
|
2025-03-14 03:28:53 +08:00
|
|
|
|
config = json.loads(str(self))
|
2025-03-29 17:33:09 +08:00
|
|
|
|
|
2025-03-14 03:28:53 +08:00
|
|
|
|
# 过滤掉 api_key
|
|
|
|
|
|
for model in config.get("custom_models", []):
|
|
|
|
|
|
model["api_key"] = DEFAULT_MOCK_API if model.get("api_key") else ""
|
|
|
|
|
|
|
|
|
|
|
|
return config
|
2025-03-29 17:33:09 +08:00
|
|
|
|
|
2025-03-14 03:28:53 +08:00
|
|
|
|
def compare_custom_models(self, value):
|
|
|
|
|
|
"""
|
|
|
|
|
|
比较 custom_models 中的 api_key,如果输入的 api_key 与当前的 api_key 相同,则不修改
|
|
|
|
|
|
如果输入的 api_key 为 DEFAULT_MOCK_API,则使用当前的 api_key
|
|
|
|
|
|
"""
|
2025-03-29 17:46:19 +08:00
|
|
|
|
current_models_dict = {model["custom_id"]: model.get("api_key") for model in self.get("custom_models", [])}
|
2025-03-14 03:28:53 +08:00
|
|
|
|
|
|
|
|
|
|
for i, model in enumerate(value):
|
|
|
|
|
|
input_custom_id = model.get("custom_id")
|
|
|
|
|
|
input_api_key = model.get("api_key")
|
|
|
|
|
|
|
|
|
|
|
|
if input_custom_id in current_models_dict:
|
|
|
|
|
|
current_api_key = current_models_dict[input_custom_id]
|
|
|
|
|
|
if input_api_key == DEFAULT_MOCK_API or input_api_key == current_api_key:
|
|
|
|
|
|
value[i]["api_key"] = current_api_key
|
|
|
|
|
|
|
|
|
|
|
|
return value
|