summaryrefslogtreecommitdiffstats
path: root/src/client/cohere.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-23 16:46:48 +0800
committerGitHub <noreply@github.com>2024-04-23 16:46:48 +0800
commitd1aafa11153ab689c21c2c57c47da52337d8e8d1 (patch)
treedc20dc033e9d376aab09941835a842b22fe32c02 /src/client/cohere.rs
parent1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff)
downloadaichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/cohere.rs')
-rw-r--r--src/client/cohere.rs10
1 files changed, 7 insertions, 3 deletions
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index a92e238..bfea105 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -67,7 +67,7 @@ impl CohereClient {
fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result<RequestBuilder> {
let api_key = self.get_api_key().ok();
- let body = build_body(data, self.model.name.clone())?;
+ let body = build_body(data, &self.model)?;
let url = API_URL;
@@ -131,7 +131,7 @@ fn check_error(data: &Value) -> Result<()> {
}
}
-pub(crate) fn build_body(data: SendData, model: String) -> Result<Value> {
+pub(crate) fn build_body(data: SendData, model: &Model) -> Result<Value> {
let SendData {
mut messages,
temperature,
@@ -179,9 +179,13 @@ pub(crate) fn build_body(data: SendData, model: String) -> Result<Value> {
let message = message["message"].as_str().unwrap_or_default();
let mut body = json!({
- "model": model,
+ "model": &model.name,
"message": message,
});
+
+ if let Some(max_tokens) = model.max_output_tokens {
+ body["max_tokens"] = max_tokens.into();
+ }
if !messages.is_empty() {
body["chat_history"] = messages.into();