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-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 ) :
2025-04-02 13:00:25 +08:00
return self . get ( self . __key ( key ) )
2024-07-07 01:58:23 +08:00
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 ( ) }
2025-05-05 19:40:23 +08:00
def update ( self , other ) :
for key , value in other . items ( ) :
self [ key ] = value
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-06-26 02:13:46 +08:00
self . filename = str ( Path ( f " { self . save_dir } /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
2025-04-06 20:33:15 +08:00
self . _update_models_from_file ( )
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-05-24 11:29:45 +08:00
self . add_item ( " enable_web_search " , default = False , des = " 是否开启网页搜索(注:现阶段会根据 TAVILY_API_KEY 自动开启,无法手动配置,将会在下个版本移除此配置项) " ) # noqa: E501
2025-05-02 23:56:59 +08:00
# 默认智能体配置
self . add_item ( " default_agent_id " , default = " " , des = " 默认智能体ID " )
2024-07-17 18:52:20 +08:00
# 模型配置
## 注意这里是模型名,而不是具体的模型路径,默认使用 HuggingFace 的路径
2025-02-20 01:26:12 +08:00
## 如果需要自定义本地模型路径,则在 src/.env 中配置 MODEL_DIR
2025-04-06 20:33:15 +08:00
self . add_item ( " model_provider " , default = " siliconflow " , des = " 模型提供商 " , choices = list ( self . 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-26 23:58:26 +08:00
2025-04-06 20:33:15 +08:00
self . add_item ( " embed_model " , default = " siliconflow/BAAI/bge-m3 " , des = " Embedding 模型 " , choices = list ( self . embed_model_names . keys ( ) ) )
2025-05-24 11:29:45 +08:00
self . add_item ( " reranker " , default = " siliconflow/BAAI/bge-reranker-v2-m3 " , des = " Re-Ranker 模型 " , choices = list ( self . reranker_names . keys ( ) ) ) # noqa: E501
2024-07-25 20:30:28 +08:00
self . add_item ( " model_local_paths " , default = { } , des = " 本地模型路径 " )
2025-05-09 23:47:45 +08:00
self . add_item ( " use_rewrite_query " , default = " on " , 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 }
2025-04-06 20:33:15 +08:00
def _update_models_from_file ( self ) :
"""
从 models . yaml 和 models . private . yml 中更新 MODEL_NAMES
"""
2025-05-24 11:29:45 +08:00
with open ( Path ( " src/static/models.yaml " ) , encoding = ' utf-8 ' ) as f :
2025-04-06 20:33:15 +08:00
_models = yaml . safe_load ( f )
# 尝试打开一个 models.private.yml 文件,用来覆盖 models.yaml 中的配置
try :
2025-05-24 11:29:45 +08:00
with open ( Path ( " src/static/models.private.yml " ) , encoding = ' utf-8 ' ) as f :
2025-04-06 20:33:15 +08:00
_models_private = yaml . safe_load ( f )
except FileNotFoundError :
_models_private = { }
2025-02-28 02:39:47 +08:00
2025-04-10 11:41:21 +08:00
# 修改为按照子元素合并
# _models = {**_models, **_models_private}
2025-04-06 20:33:15 +08:00
2025-04-13 23:12:55 +08:00
self . model_names = { * * _models [ " MODEL_NAMES " ] , * * _models_private . get ( " MODEL_NAMES " , { } ) }
self . embed_model_names = { * * _models [ " EMBED_MODEL_INFO " ] , * * _models_private . get ( " EMBED_MODEL_INFO " , { } ) }
self . reranker_names = { * * _models [ " RERANKER_LIST " ] , * * _models_private . get ( " RERANKER_LIST " , { } ) }
2025-04-06 20:33:15 +08:00
def _save_models_to_file ( self ) :
_models = {
" MODEL_NAMES " : self . model_names ,
" EMBED_MODEL_INFO " : self . embed_model_names ,
" RERANKER_LIST " : self . reranker_names ,
}
with open ( Path ( " src/static/models.private.yml " ) , ' w ' , encoding = ' utf-8 ' ) as f :
yaml . dump ( _models , f , indent = 2 , allow_unicode = True )
def handle_self ( self ) :
"""
处理配置
"""
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-04-13 23:50:19 +08:00
if self . model_dir :
if os . path . exists ( self . model_dir ) :
2025-06-24 00:14:12 +08:00
logger . debug ( f " The model directory ( { self . model_dir } ) contains the following folders: { os . listdir ( self . model_dir ) } " )
2025-04-13 23:50:19 +08:00
else :
2025-06-24 00:14:12 +08:00
logger . warning ( f " Warning: The model directory ( { self . model_dir } ) does not exist. If not configured, please ignore it. If configured, please check if the configuration is correct;"
" For example, the mapping in the docker-compose file " )
2025-04-13 23:50:19 +08:00
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的环境变量
2025-04-09 12:05:05 +08:00
# 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
# 2025.04.08 修改为不手动配置, 只要配置了TAVILY_API_KEY, 就默认开启web_search
if os . getenv ( " TAVILY_API_KEY " ) :
self . enable_web_search = True
2025-02-20 01:26:12 +08:00
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 " ) :
2025-05-24 11:29:45 +08:00
with open ( self . filename ) 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 " ) :
2025-05-24 11:29:45 +08:00
with open ( self . filename ) 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
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
2025-05-05 16:31:35 +08:00
def dump_config ( self ) :
return json . loads ( str ( self ) )
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