summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-23 18:14:47 +0800
committerGitHub <noreply@github.com>2024-04-23 18:14:47 +0800
commit9c6c9f10a27d0993636b453f39d8934c95c5c2b2 (patch)
tree268598294b234c7eac87fc70d1e2d1b32ac88a78 /src/client/cohere.rs
parentd1aafa11153ab689c21c2c57c47da52337d8e8d1 (diff)
downloadaichat-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.rs19
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();
}