diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-23 16:46:48 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-23 16:46:48 +0800 |
| commit | d1aafa11153ab689c21c2c57c47da52337d8e8d1 (patch) | |
| tree | dc20dc033e9d376aab09941835a842b22fe32c02 /src/client/cohere.rs | |
| parent | 1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff) | |
| download | aichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz | |
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/cohere.rs')
| -rw-r--r-- | src/client/cohere.rs | 10 |
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(); |
