1414from common .utils .utils import prepare_model_arg
1515from langchain_community .llms import VLLMOpenAI
1616from langchain_openai import AzureChatOpenAI
17+
18+
1719# from langchain_community.llms import Tongyi, VLLM
1820
1921class LLMConfig (BaseModel ):
@@ -24,16 +26,17 @@ class LLMConfig(BaseModel):
2426 api_key : Optional [str ] = None
2527 api_base_url : Optional [str ] = None
2628 additional_params : Dict [str , Any ] = {}
29+
2730 class Config :
2831 frozen = True
2932
3033 def __hash__ (self ):
3134 if hasattr (self , 'additional_params' ) and isinstance (self .additional_params , dict ):
32- hashable_params = frozenset ((k , tuple (v ) if isinstance (v , (list , dict )) else v )
33- for k , v in self .additional_params .items ())
35+ hashable_params = frozenset ((k , tuple (v ) if isinstance (v , (list , dict )) else v )
36+ for k , v in self .additional_params .items ())
3437 else :
3538 hashable_params = None
36-
39+
3740 return hash ((
3841 self .model_id ,
3942 self .model_type ,
@@ -61,6 +64,7 @@ def llm(self) -> BaseChatModel:
6164 """Return the langchain LLM instance"""
6265 return self ._llm
6366
67+
6468class OpenAIvLLM (BaseLLM ):
6569 def _init_llm (self ) -> VLLMOpenAI :
6670 return VLLMOpenAI (
@@ -71,6 +75,7 @@ def _init_llm(self) -> VLLMOpenAI:
7175 ** self .config .additional_params ,
7276 )
7377
78+
7479class OpenAIAzureLLM (BaseLLM ):
7580 def _init_llm (self ) -> AzureChatOpenAI :
7681 api_version = self .config .additional_params .get ("api_version" )
@@ -88,6 +93,8 @@ def _init_llm(self) -> AzureChatOpenAI:
8893 streaming = True ,
8994 ** self .config .additional_params ,
9095 )
96+
97+
9198class OpenAILLM (BaseLLM ):
9299 def _init_llm (self ) -> BaseChatModel :
93100 return BaseChatOpenAI (
@@ -138,26 +145,30 @@ def register_llm(cls, model_type: str, llm_class: Type[BaseLLM]):
138145 return config """
139146
140147
141- async def get_default_config () -> LLMConfig :
148+ async def get_default_config (custom_model_id : Optional [ int ] = None ) -> LLMConfig :
142149 with Session (engine ) as session :
143- db_model = session .exec (
144- select (AiModelDetail ).where (AiModelDetail .default_model == True )
145- ).first ()
150+ db_model : AiModelDetail | None = None
151+ if custom_model_id :
152+ db_model = session .get (AiModelDetail , custom_model_id )
153+ if not db_model :
154+ db_model = session .exec (
155+ select (AiModelDetail ).where (AiModelDetail .default_model == True )
156+ ).first ()
146157 if not db_model :
147158 raise Exception ("The system default model has not been set" )
148159
149160 additional_params = {}
150161 if db_model .config :
151162 try :
152163 config_raw = json .loads (db_model .config )
153- additional_params = {item ["key" ]: prepare_model_arg (item .get ('val' )) for item in config_raw if "key" in item and "val" in item }
164+ additional_params = {item ["key" ]: prepare_model_arg (item .get ('val' )) for item in config_raw if
165+ "key" in item and "val" in item }
154166 except Exception :
155167 pass
156168 if not db_model .api_domain .startswith ("http" ):
157169 db_model .api_domain = await sqlbot_decrypt (db_model .api_domain )
158170 if db_model .api_key :
159171 db_model .api_key = await sqlbot_decrypt (db_model .api_key )
160-
161172
162173 # 构造 LLMConfig
163174 return LLMConfig (
0 commit comments