diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-23 18:14:47 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-23 18:14:47 +0800 |
| commit | 9c6c9f10a27d0993636b453f39d8934c95c5c2b2 (patch) | |
| tree | 268598294b234c7eac87fc70d1e2d1b32ac88a78 /src/client/cohere.rs | |
| parent | d1aafa11153ab689c21c2c57c47da52337d8e8d1 (diff) | |
| download | aichat-9c6c9f10a27d0993636b453f39d8934c95c5c2b2.tar.gz | |
feat: builtin models can be overwrited by models config (#429)
Diffstat (limited to 'src/client/cohere.rs')
| -rw-r--r-- | src/client/cohere.rs | 19 |
1 files changed, 5 insertions, 14 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs index bfea105..445c145 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -1,6 +1,6 @@ use super::{ json_stream, message::*, patch_system_message, Client, CohereClient, ExtraConfig, Model, - PromptType, ReplyHandler, SendData, + ModelConfig, PromptType, ReplyHandler, SendData, }; use crate::utils::PromptKind; @@ -23,6 +23,8 @@ const MODELS: [(&str, usize, &str); 2] = [ pub struct CohereConfig { pub name: Option<String>, pub api_key: Option<String>, + #[serde(default)] + pub models: Vec<ModelConfig>, pub extra: Option<ExtraConfig>, } @@ -47,23 +49,12 @@ impl Client for CohereClient { } impl CohereClient { + list_models_fn!(CohereConfig, &MODELS); config_get_fn!(api_key, get_api_key); pub const PROMPTS: [PromptType<'static>; 1] = [("api_key", "API Key:", false, PromptKind::String)]; - pub fn list_models(local_config: &CohereConfig) -> Vec<Model> { - let client_name = Self::name(local_config); - MODELS - .into_iter() - .map(|(name, max_input_tokens, capabilities)| { - Model::new(client_name, name) - .set_capabilities(capabilities.into()) - .set_max_input_tokens(Some(max_input_tokens)) - }) - .collect() - } - fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> { let api_key = self.get_api_key().ok(); @@ -182,7 +173,7 @@ pub(crate) fn build_body(data: SendData, model: &Model) -> Result<Value> { "model": &model.name, "message": message, }); - + if let Some(max_tokens) = model.max_output_tokens { body["max_tokens"] = max_tokens.into(); } |
