6060_storage_class_cache : dict [str , 'type[AsyncBaseDb]' ] = {}
6161
6262
63- def _resolve_provider_class (provider : str ) -> 'type[Model] | None' :
64- """
65- 按需加载并缓存提供商模型类
66-
67- :param provider: 提供商名称
68- :return: 模型类,未找到返回None
63+ class AiUtil :
6964 """
70- if provider in _provider_class_cache :
71- return _provider_class_cache [provider ]
72- entry = _PROVIDER_REGISTRY .get (provider )
73- if entry is None :
74- return None
75- module_path , class_name = entry
76- cls = getattr (import_module (module_path ), class_name )
77- _provider_class_cache [provider ] = cls
78- return cls
79-
80-
81- def _resolve_storage_class (db_type : str ) -> 'type[AsyncBaseDb]' :
65+ AI工具类
8266 """
83- 按需加载并缓存存储引擎类
8467
85- :param db_type: 数据库类型
86- :return: 存储引擎类
87- """
88- if db_type in _storage_class_cache :
89- return _storage_class_cache [db_type ]
90- entry = _STORAGE_ENGINE_REGISTRY .get (db_type )
91- if entry is None :
92- # 默认使用MySQL
93- entry = _STORAGE_ENGINE_REGISTRY ['mysql' ]
94- module_path , class_name = entry
95- cls = getattr (import_module (module_path ), class_name )
96- _storage_class_cache [db_type ] = cls
97- return cls
68+ @classmethod
69+ def _resolve_provider_class (cls , provider : str ) -> 'type[Model] | None' :
70+ """
71+ 按需加载并缓存提供商模型类
72+
73+ :param provider: 提供商名称
74+ :return: 模型类,未找到返回None
75+ """
76+ if provider in _provider_class_cache :
77+ return _provider_class_cache [provider ]
78+ entry = _PROVIDER_REGISTRY .get (provider )
79+ if entry is None :
80+ return None
81+ module_path , class_name = entry
82+ provider_cls = getattr (import_module (module_path ), class_name )
83+ _provider_class_cache [provider ] = provider_cls
84+ return provider_cls
9885
86+ @classmethod
87+ def _resolve_storage_class (cls , db_type : str ) -> 'type[AsyncBaseDb]' :
88+ """
89+ 按需加载并缓存存储引擎类
9990
100- class AiUtil :
101- """
102- AI工具类
103- """
91+ :param db_type: 数据库类型
92+ :return: 存储引擎类
93+ """
94+ if db_type in _storage_class_cache :
95+ return _storage_class_cache [db_type ]
96+ entry = _STORAGE_ENGINE_REGISTRY .get (db_type )
97+ if entry is None :
98+ # 默认使用MySQL
99+ entry = _STORAGE_ENGINE_REGISTRY ['mysql' ]
100+ module_path , class_name = entry
101+ storage_cls = getattr (import_module (module_path ), class_name )
102+ _storage_class_cache [db_type ] = storage_cls
103+ return storage_cls
104104
105105 @classmethod
106106 def get_storage_engine (cls ) -> 'AsyncBaseDb' :
@@ -109,7 +109,7 @@ def get_storage_engine(cls) -> 'AsyncBaseDb':
109109
110110 :return: 存储引擎实例
111111 """
112- storage_engine_class = _resolve_storage_class (DataBaseConfig .db_type )
112+ storage_engine_class = cls . _resolve_storage_class (DataBaseConfig .db_type )
113113
114114 return storage_engine_class (
115115 db_engine = async_engine ,
@@ -164,9 +164,9 @@ def get_model_from_factory(
164164 params ['host' ] = base_url
165165 if provider == 'DashScope' and not base_url :
166166 params ['base_url' ] = 'https://dashscope.aliyuncs.com/compatible-mode/v1'
167- model_class = _resolve_provider_class (provider )
167+ model_class = cls . _resolve_provider_class (provider )
168168 if model_class is None :
169169 # 未知提供商,回退到OpenAI
170- model_class = _resolve_provider_class ('OpenAI' )
170+ model_class = cls . _resolve_provider_class ('OpenAI' )
171171
172172 return model_class (** params )
0 commit comments