diff options
| author | sigoden <sigoden@gmail.com> | 2024-12-04 21:03:59 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-12-04 21:03:59 +0800 |
| commit | 3a3388375be05758d5f5574cce1d0f8eb7bdd604 (patch) | |
| tree | cc8e19e218c20f28a04feca33850fce4bb2cbe98 /src/config | |
| parent | 7d42fe9429f75d195f865b07cef10d040d5397f2 (diff) | |
| download | aichat-3a3388375be05758d5f5574cce1d0f8eb7bdd604.tar.gz | |
refactor: improve retrieve model (#1036)
- check the model type while retrieve model
- select chat/reranker model even if it is missed in client models
- find predefined-models for openai-compatible client with startsWith
- remove client::ApiType
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/agent.rs | 2 | ||||
| -rw-r--r-- | src/config/mod.rs | 19 | ||||
| -rw-r--r-- | src/config/session.rs | 2 |
3 files changed, 13 insertions, 10 deletions
diff --git a/src/config/agent.rs b/src/config/agent.rs index 5e54b26..c81ed11 100644 --- a/src/config/agent.rs +++ b/src/config/agent.rs @@ -61,7 +61,7 @@ impl Agent { let model = { let config = config.read(); match agent_config.model_id.as_ref() { - Some(model_id) => Model::retrieve_chat(&config, model_id)?, + Some(model_id) => Model::retrieve_model(&config, model_id, ModelType::Chat)?, None => config.current_model().clone(), } }; diff --git a/src/config/mod.rs b/src/config/mod.rs index db210ab..4c6a321 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -11,8 +11,8 @@ pub use self::role::{ use self::session::Session; use crate::client::{ - create_client_config, list_chat_models, list_client_types, list_reranker_models, ClientConfig, - MessageContentToolCalls, Model, OPENAI_COMPATIBLE_PLATFORMS, + create_client_config, list_client_types, list_models, ClientConfig, MessageContentToolCalls, + Model, ModelType, OPENAI_COMPATIBLE_PLATFORMS, }; use crate::function::{FunctionDeclaration, Functions, ToolResult}; use crate::rag::Rag; @@ -775,7 +775,7 @@ impl Config { pub fn set_rag_reranker_model(config: &GlobalConfig, value: Option<String>) -> Result<()> { if let Some(id) = &value { - Model::retrieve_reranker(&config.read(), id)?; + Model::retrieve_model(&config.read(), id, ModelType::Reranker)?; } let has_rag = config.read().rag.is_some(); match has_rag { @@ -822,7 +822,7 @@ impl Config { } pub fn set_model(&mut self, model_id: &str) -> Result<()> { - let model = Model::retrieve_chat(self, model_id)?; + let model = Model::retrieve_model(self, model_id, ModelType::Chat)?; match self.role_like_mut() { Some(role_like) => role_like.set_model(&model), None => { @@ -893,7 +893,7 @@ impl Config { match role.model_id() { Some(model_id) => { if self.model.id() != model_id { - let model = Model::retrieve_chat(self, model_id)?; + let model = Model::retrieve_model(self, model_id, ModelType::Chat)?; role.set_model(&model); } else { role.set_model(&self.model); @@ -1666,7 +1666,7 @@ impl Config { if args.len() == 1 { values = match cmd { ".role" => map_completion_values(Self::list_roles(true)), - ".model" => list_chat_models(self) + ".model" => list_models(self, ModelType::Chat) .into_iter() .map(|v| (v.id(), Some(v.description()))) .collect(), @@ -1761,7 +1761,10 @@ impl Config { }; complete_option_bool(save_session) } - "rag_reranker_model" => list_reranker_models(self).iter().map(|v| v.id()).collect(), + "rag_reranker_model" => list_models(self, ModelType::Reranker) + .iter() + .map(|v| v.id()) + .collect(), "highlight" => complete_bool(self.highlight), _ => vec![], }; @@ -2268,7 +2271,7 @@ impl Config { fn setup_model(&mut self) -> Result<()> { let mut model_id = self.model_id.clone(); if model_id.is_empty() { - let models = list_chat_models(self); + let models = list_models(self, ModelType::Chat); if models.is_empty() { bail!("No available model"); } diff --git a/src/config/session.rs b/src/config/session.rs index 820fa12..227a23c 100644 --- a/src/config/session.rs +++ b/src/config/session.rs @@ -82,7 +82,7 @@ impl Session { let mut session: Self = serde_yaml::from_str(&content).with_context(|| format!("Invalid session {}", name))?; - session.model = Model::retrieve_chat(config, &session.model_id)?; + session.model = Model::retrieve_model(config, &session.model_id, ModelType::Chat)?; if let Some(autoname) = name.strip_prefix("_/") { session.name = TEMP_SESSION_NAME.to_string(); |
