diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-14 12:43:16 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-14 12:43:16 +0800 |
| commit | 79d0bba640d954cd3e6acd7f4e83900eb9d56a1c (patch) | |
| tree | e266ca283735b91bc2964ec4040b088cf7b865ed /src/client/common.rs | |
| parent | 5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (diff) | |
| download | aichat-79d0bba640d954cd3e6acd7f4e83900eb9d56a1c.tar.gz | |
feat: allow binding model to the role (#505)
Diffstat (limited to 'src/client/common.rs')
| -rw-r--r-- | src/client/common.rs | 12 |
1 files changed, 6 insertions, 6 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 495160b..004a7dc 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -70,8 +70,7 @@ macro_rules! register_client { impl $client { pub const NAME: &'static str = $name; - pub fn init(global_config: &$crate::config::GlobalConfig) -> Option<Box<dyn Client>> { - let model = global_config.read().model.clone(); + pub fn init(global_config: &$crate::config::GlobalConfig, model: &$crate::client::Model) -> Option<Box<dyn Client>> { let config = global_config.read().clients.iter().find_map(|client_config| { if let ClientConfig::$config(c) = client_config { if Self::name(c) == &model.client_name { @@ -84,7 +83,7 @@ macro_rules! register_client { Some(Box::new(Self { global_config: global_config.clone(), config, - model, + model: model.clone(), })) } @@ -109,11 +108,12 @@ macro_rules! register_client { )+ - pub fn init_client(config: &$crate::config::GlobalConfig) -> anyhow::Result<Box<dyn Client>> { + pub fn init_client(config: &$crate::config::GlobalConfig, model: Option<$crate::client::Model>) -> anyhow::Result<Box<dyn Client>> { + let model = model.unwrap_or_else(|| config.read().model.clone()); None - $(.or_else(|| $client::init(config)))+ + $(.or_else(|| $client::init(config, &model)))+ .ok_or_else(|| { - anyhow::anyhow!("Unknown client '{}'", &config.read().model.client_name) + anyhow::anyhow!("Unknown client '{}'", model.client_name) }) } |
